sleep_ai_demo / db /src /sleep_db /query.py
lokework's picture
Upload sleep full-agent demo
6025aa5 verified
Raw
History Blame Contribute Delete
2.68 kB
from __future__ import annotations
import argparse
import json
from typing import Any
from sleep_db.config import load_settings
from sleep_db.retrieval import retrieve_domain_knowledge, retrieve_knowledge
from sleep_db.seed_sqlite import get_sleep_metrics
def build_query_response(
query: str,
user_id: str,
date: str,
top_k: int = 5,
) -> dict[str, Any]:
settings = load_settings()
metrics = get_sleep_metrics(user_id, date, settings.sqlite_path)
retrieval = retrieve_knowledge(
query=query,
user_metrics=metrics,
baseline_metrics={
"heart_rate_baseline": metrics["heart_rate_baseline"],
"hrv_baseline": metrics["hrv_baseline"],
},
top_k=top_k,
settings=settings,
)
return {
"query": query,
"sqlite_metrics": metrics,
"neo4j_graph_results": retrieval["graph_results"],
"milvus_vector_results": retrieval["vector_results"],
"merged_context": retrieval["merged_context"],
"sources": retrieval["sources"],
}
def build_domain_query_response(
query: str,
user_id: str,
date: str,
domains: list[str],
top_k: int = 5,
) -> dict[str, Any]:
settings = load_settings()
metrics = get_sleep_metrics(user_id, date, settings.sqlite_path)
retrieval = retrieve_domain_knowledge(
query=query,
user_metrics=metrics,
domains=domains,
baseline_metrics={
"heart_rate_baseline": metrics["heart_rate_baseline"],
"hrv_baseline": metrics["hrv_baseline"],
},
top_k=top_k,
settings=settings,
)
return {
"query": query,
"requested_domains": retrieval["domains"],
"sqlite_metrics": metrics,
"neo4j_graph_results": retrieval["graph_results"],
"milvus_vector_results": retrieval["vector_results"],
"merged_context": retrieval["merged_context"],
"sources": retrieval["sources"],
}
def main() -> None:
parser = argparse.ArgumentParser(
description="Run a query and show what each sleep DB layer retrieved."
)
parser.add_argument("query", help="Natural-language query to retrieve against.")
parser.add_argument("--user-id", default="demo_user")
parser.add_argument("--date", default="2026-05-30")
parser.add_argument("--top-k", type=int, default=5)
args = parser.parse_args()
response = build_query_response(
query=args.query,
user_id=args.user_id,
date=args.date,
top_k=args.top_k,
)
print(json.dumps(response, indent=2, ensure_ascii=False))
if __name__ == "__main__":
main()