build_creation_demo.py 3.0 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495
  1. #!/usr/bin/env python3
  2. """Build formal creation-query batches.
  3. Default mode only prints a summary. Use --export-json for a local review file
  4. or --persist to write query_batches/queries into the formal PostgreSQL store.
  5. """
  6. from __future__ import annotations
  7. import argparse
  8. import json
  9. from pathlib import Path
  10. from acquisition.queries.builder import (
  11. QueryBuildOptions,
  12. TREES,
  13. build_creation_query_batch,
  14. persist_query_batch,
  15. )
  16. from acquisition.repositories.postgres import PostgresAcquisitionRepository
  17. from core.config import CreationDbConfig, Settings
  18. from core.db_session import transaction
  19. def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
  20. parser = argparse.ArgumentParser(description=__doc__)
  21. parser.add_argument("--env-file", default=".env")
  22. parser.add_argument("--tree-path", type=Path)
  23. parser.add_argument("--per", type=int, default=30)
  24. parser.add_argument("--batch-n", type=int, default=30)
  25. parser.add_argument("--seed", type=int, default=7)
  26. parser.add_argument("--dry", action="store_true", help="Skip LLM query filtering")
  27. parser.add_argument("--export-json", type=Path, help="Optional local review export")
  28. parser.add_argument("--persist", action="store_true", help="Persist to formal PG")
  29. parser.add_argument("--name", default="creation-demo")
  30. return parser.parse_args(argv)
  31. def _summary(generated: dict) -> dict:
  32. families = generated.get("families") or []
  33. total = sum(len(f.get("items") or []) for f in families)
  34. kept = sum(
  35. 1
  36. for family in families
  37. for item in family.get("items") or []
  38. if item.get("keep", True)
  39. )
  40. return {
  41. "family_count": len(families),
  42. "query_count": total,
  43. "kept_count": kept,
  44. "metadata": generated.get("metadata") or {},
  45. }
  46. def main(argv: list[str] | None = None) -> int:
  47. args = parse_args(argv)
  48. settings = Settings.from_env(args.env_file)
  49. generated = build_creation_query_batch(
  50. settings,
  51. tree_path=args.tree_path or TREES,
  52. options=QueryBuildOptions(
  53. per=args.per,
  54. batch_n=args.batch_n,
  55. seed=args.seed,
  56. dry=args.dry,
  57. ),
  58. )
  59. summary = _summary(generated)
  60. if args.export_json:
  61. args.export_json.parent.mkdir(parents=True, exist_ok=True)
  62. args.export_json.write_text(
  63. json.dumps(generated, ensure_ascii=False, indent=1),
  64. encoding="utf-8",
  65. )
  66. summary["export_json"] = str(args.export_json)
  67. if args.persist:
  68. db_config = CreationDbConfig.from_env(args.env_file)
  69. with transaction(db_config) as conn:
  70. repo = PostgresAcquisitionRepository(conn)
  71. batch, count = persist_query_batch(
  72. repo,
  73. generated,
  74. name=args.name,
  75. )
  76. summary["batch_id"] = str(batch.id)
  77. summary["persisted_queries"] = count
  78. print(json.dumps(summary, ensure_ascii=False, indent=2))
  79. return 0
  80. if __name__ == "__main__":
  81. raise SystemExit(main())