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()