run_creation_pipeline.py 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122
  1. #!/usr/bin/env python3
  2. """Run formal acquisition, decode every coarse-hit item, and build payload drafts."""
  3. from __future__ import annotations
  4. import argparse
  5. import json
  6. from dataclasses import asdict, is_dataclass
  7. from uuid import UUID
  8. from acquisition.repositories.postgres import PostgresAcquisitionRepository
  9. from acquisition.runner import DEFAULT_PLATFORMS, run_batch
  10. from core.config import CreationDbConfig, Settings
  11. from core.db_session import transaction
  12. from decode_content.repositories.postgres import PostgresDecodeRepository
  13. from decode_content.service import DecodeService
  14. from pipeline.decode_runner import run_decode_stage
  15. def _model_dump(value):
  16. if hasattr(value, "model_dump"):
  17. return value.model_dump(mode="json")
  18. if is_dataclass(value):
  19. return asdict(value)
  20. if isinstance(value, dict):
  21. return value
  22. return dict(value)
  23. def _dry_ingest_payloads(repo: PostgresDecodeRepository, outputs: list) -> list[dict]:
  24. records: list[dict] = []
  25. for output in outputs:
  26. for draft in output.payload_drafts:
  27. if draft.id is None:
  28. continue
  29. repo.mark_payload_draft_ingested(draft.id)
  30. record = repo.save_ingest_record(
  31. payload_draft_id=draft.id,
  32. target_system="dry-run",
  33. target_id=str(draft.id),
  34. status="ingested",
  35. response_payload={
  36. "dry_run": True,
  37. "note": "payload generated by formal creation pipeline; external ingest API not called",
  38. "payload": draft.payload,
  39. },
  40. )
  41. records.append(_model_dump(record))
  42. return records
  43. def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
  44. parser = argparse.ArgumentParser(description=__doc__)
  45. parser.add_argument("--batch-id", required=True, help="Formal query batch UUID")
  46. parser.add_argument(
  47. "--platform",
  48. action="append",
  49. choices=DEFAULT_PLATFORMS,
  50. help="Platform to run. Repeat to run multiple platforms. Default: all.",
  51. )
  52. parser.add_argument("--search-limit", type=int, default=10)
  53. parser.add_argument("--display-limit", type=int, default=5)
  54. parser.add_argument("--decode-limit", type=int, default=100)
  55. parser.add_argument("--run-key")
  56. parser.add_argument("--env-file", default=".env")
  57. parser.add_argument("--no-resume", action="store_true")
  58. parser.add_argument("--no-skip-done", action="store_true")
  59. parser.add_argument("--no-dry-ingest-record", action="store_true")
  60. return parser.parse_args(argv)
  61. def main(argv: list[str] | None = None) -> int:
  62. args = parse_args(argv)
  63. settings = Settings.from_env(args.env_file)
  64. db_config = CreationDbConfig.from_env(args.env_file)
  65. platforms = tuple(args.platform or DEFAULT_PLATFORMS)
  66. batch_id = UUID(args.batch_id)
  67. with transaction(db_config) as conn:
  68. acquisition_repo = PostgresAcquisitionRepository(conn)
  69. acquisition = run_batch(
  70. acquisition_repo,
  71. batch_id=batch_id,
  72. settings=settings,
  73. platforms=platforms,
  74. search_limit=args.search_limit,
  75. display_limit=args.display_limit,
  76. classify=True,
  77. resume=not args.no_resume,
  78. skip_done=not args.no_skip_done,
  79. run_key=args.run_key,
  80. )
  81. with transaction(db_config) as conn:
  82. acquisition_repo = PostgresAcquisitionRepository(conn)
  83. decode_repo = PostgresDecodeRepository(conn)
  84. decode_service = DecodeService(settings=settings, repository=decode_repo)
  85. decode = run_decode_stage(
  86. candidate_repo=acquisition_repo,
  87. decode_service=decode_service,
  88. run_id=acquisition.run_id,
  89. limit=args.decode_limit,
  90. )
  91. ingest_records = [] if args.no_dry_ingest_record else _dry_ingest_payloads(decode_repo, decode.outputs)
  92. print(json.dumps(
  93. {
  94. "batch_id": str(batch_id),
  95. "run_id": str(acquisition.run_id),
  96. "platforms": list(platforms),
  97. "acquisition": _model_dump(acquisition),
  98. "decode": _model_dump(decode),
  99. "dry_ingest_records": ingest_records,
  100. },
  101. ensure_ascii=False,
  102. default=str,
  103. indent=2,
  104. ))
  105. return 0
  106. if __name__ == "__main__":
  107. raise SystemExit(main())