run_creation_pipeline.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120
  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, IngestApiConfig, Settings
  11. from core.db_session import transaction
  12. from decode_content.ingest import KnowledgeIngestClient, ingest_payload_draft
  13. from decode_content.repositories.postgres import PostgresDecodeRepository
  14. from decode_content.service import DecodeService
  15. from pipeline.decode_runner import run_decode_stage
  16. def _model_dump(value):
  17. if hasattr(value, "model_dump"):
  18. return value.model_dump(mode="json")
  19. if is_dataclass(value):
  20. return asdict(value)
  21. if isinstance(value, dict):
  22. return value
  23. return dict(value)
  24. def _ingest_payloads(repo: PostgresDecodeRepository, outputs: list, *, dry_run: bool, env_file: str) -> list[dict]:
  25. records: list[dict] = []
  26. client = None if dry_run else KnowledgeIngestClient(IngestApiConfig.from_env(env_file))
  27. for output in outputs:
  28. for draft in output.payload_drafts:
  29. if draft.id is None:
  30. continue
  31. record = ingest_payload_draft(repo, draft, dry_run=dry_run, client=client)
  32. records.append(_model_dump(record))
  33. return records
  34. def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
  35. parser = argparse.ArgumentParser(description=__doc__)
  36. parser.add_argument("--batch-id", required=True, help="Formal query batch UUID")
  37. parser.add_argument(
  38. "--platform",
  39. action="append",
  40. choices=DEFAULT_PLATFORMS,
  41. help="Platform to run. Repeat to run multiple platforms. Default: all.",
  42. )
  43. parser.add_argument("--search-limit", type=int, default=10)
  44. parser.add_argument("--display-limit", type=int, default=5)
  45. parser.add_argument("--decode-limit", type=int, default=100)
  46. parser.add_argument("--run-key")
  47. parser.add_argument("--env-file", default=".env")
  48. parser.add_argument("--no-resume", action="store_true")
  49. parser.add_argument("--no-skip-done", action="store_true")
  50. parser.add_argument("--no-dry-ingest-record", action="store_true")
  51. parser.add_argument("--real-ingest", action="store_true", help="Call CK_INGEST_API_URL instead of dry-run records.")
  52. return parser.parse_args(argv)
  53. def main(argv: list[str] | None = None) -> int:
  54. args = parse_args(argv)
  55. settings = Settings.from_env(args.env_file)
  56. db_config = CreationDbConfig.from_env(args.env_file)
  57. platforms = tuple(args.platform or DEFAULT_PLATFORMS)
  58. batch_id = UUID(args.batch_id)
  59. with transaction(db_config) as conn:
  60. acquisition_repo = PostgresAcquisitionRepository(conn)
  61. acquisition = run_batch(
  62. acquisition_repo,
  63. batch_id=batch_id,
  64. settings=settings,
  65. platforms=platforms,
  66. search_limit=args.search_limit,
  67. display_limit=args.display_limit,
  68. classify=True,
  69. resume=not args.no_resume,
  70. skip_done=not args.no_skip_done,
  71. run_key=args.run_key,
  72. )
  73. with transaction(db_config) as conn:
  74. acquisition_repo = PostgresAcquisitionRepository(conn)
  75. decode_repo = PostgresDecodeRepository(conn)
  76. decode_service = DecodeService(settings=settings, repository=decode_repo)
  77. decode = run_decode_stage(
  78. candidate_repo=acquisition_repo,
  79. decode_service=decode_service,
  80. run_id=acquisition.run_id,
  81. limit=args.decode_limit,
  82. )
  83. ingest_records = [] if args.no_dry_ingest_record else _ingest_payloads(
  84. decode_repo,
  85. decode.outputs,
  86. dry_run=not args.real_ingest,
  87. env_file=args.env_file,
  88. )
  89. print(json.dumps(
  90. {
  91. "batch_id": str(batch_id),
  92. "run_id": str(acquisition.run_id),
  93. "platforms": list(platforms),
  94. "acquisition": _model_dump(acquisition),
  95. "decode": _model_dump(decode),
  96. "ingest_records": ingest_records,
  97. "dry_ingest_records": ingest_records if not args.real_ingest else [],
  98. },
  99. ensure_ascii=False,
  100. default=str,
  101. indent=2,
  102. ))
  103. return 0
  104. if __name__ == "__main__":
  105. raise SystemExit(main())