run_creation_singleton.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177
  1. #!/usr/bin/env python3
  2. """Run one real f1/f2 creation-knowledge pipeline sample to dry-run ingest."""
  3. from __future__ import annotations
  4. import argparse
  5. import json
  6. from dataclasses import asdict, is_dataclass
  7. from datetime import datetime
  8. from pathlib import Path
  9. import sys
  10. from typing import Any
  11. ROOT = Path(__file__).resolve().parents[1]
  12. if str(ROOT) not in sys.path:
  13. sys.path.insert(0, str(ROOT))
  14. from acquisition.queries.builder import ( # noqa: E402
  15. QueryBuildOptions,
  16. TREES,
  17. build_creation_query_batch,
  18. persist_query_batch,
  19. )
  20. from acquisition.repositories.postgres import PostgresAcquisitionRepository # noqa: E402
  21. from acquisition.runner import DEFAULT_PLATFORMS, run_batch # noqa: E402
  22. from core.config import CreationDbConfig, Settings # noqa: E402
  23. from core.db_session import transaction # noqa: E402
  24. from decode_content.ingest import ingest_payload_draft # noqa: E402
  25. from decode_content.repositories.postgres import PostgresDecodeRepository # noqa: E402
  26. from decode_content.service import DecodeService # noqa: E402
  27. from pipeline.decode_runner import run_decode_stage # noqa: E402
  28. def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
  29. parser = argparse.ArgumentParser(description=__doc__)
  30. parser.add_argument("--env-file", default=".env")
  31. parser.add_argument("--tree-path", type=Path, default=TREES)
  32. parser.add_argument("--per", type=int, default=1, help="Queries per active family")
  33. parser.add_argument("--batch-n", type=int, default=30)
  34. parser.add_argument("--seed", type=int, default=7)
  35. parser.add_argument(
  36. "--platform",
  37. action="append",
  38. choices=DEFAULT_PLATFORMS,
  39. help="Platform to run. Repeat for multiple platforms. Default: xiaohongshu.",
  40. )
  41. parser.add_argument("--search-limit", type=int, default=1)
  42. parser.add_argument("--display-limit", type=int, default=1)
  43. parser.add_argument("--decode-limit", type=int, default=100)
  44. parser.add_argument("--name", default="")
  45. parser.add_argument("--frontend-base", default="http://127.0.0.1:5180/app/")
  46. return parser.parse_args(argv)
  47. def _now_key() -> str:
  48. return datetime.now().strftime("%Y%m%d-%H%M%S")
  49. def _model_dump(value: Any) -> dict[str, Any]:
  50. if hasattr(value, "model_dump"):
  51. return value.model_dump(mode="json")
  52. if is_dataclass(value):
  53. return asdict(value)
  54. if isinstance(value, dict):
  55. return value
  56. return dict(value)
  57. def _dry_ingest_payloads(
  58. repo: PostgresDecodeRepository,
  59. outputs: list[Any],
  60. ) -> list[dict[str, Any]]:
  61. records: list[dict[str, Any]] = []
  62. for output in outputs:
  63. for draft in output.payload_drafts:
  64. if draft.id is None:
  65. continue
  66. record = ingest_payload_draft(repo, draft, dry_run=True)
  67. records.append(_model_dump(record))
  68. return records
  69. def main(argv: list[str] | None = None) -> int:
  70. args = parse_args(argv)
  71. settings = Settings.from_env(args.env_file)
  72. db_config = CreationDbConfig.from_env(args.env_file)
  73. platforms = tuple(args.platform or ("xiaohongshu",))
  74. run_key_suffix = _now_key()
  75. name = args.name or f"creation-singleton-{run_key_suffix}"
  76. generated = build_creation_query_batch(
  77. settings,
  78. tree_path=args.tree_path,
  79. options=QueryBuildOptions(
  80. per=args.per,
  81. batch_n=args.batch_n,
  82. seed=args.seed,
  83. active_family_keys=("f1", "f2"),
  84. ),
  85. )
  86. query_count = sum(len(family.get("items") or []) for family in generated.get("families") or [])
  87. kept_count = sum(
  88. 1
  89. for family in generated.get("families") or []
  90. for item in family.get("items") or []
  91. if item.get("keep", True)
  92. )
  93. with transaction(db_config) as conn:
  94. acquisition_repo = PostgresAcquisitionRepository(conn)
  95. batch, persisted_count = persist_query_batch(
  96. acquisition_repo,
  97. generated,
  98. name=name,
  99. source_type="generated",
  100. generation_method="creation_singleton_v1",
  101. target_platforms=list(platforms),
  102. )
  103. with transaction(db_config) as conn:
  104. acquisition_repo = PostgresAcquisitionRepository(conn)
  105. acquisition = run_batch(
  106. acquisition_repo,
  107. batch_id=batch.id,
  108. settings=settings,
  109. platforms=platforms,
  110. search_limit=args.search_limit,
  111. display_limit=args.display_limit,
  112. classify=True,
  113. resume=False,
  114. skip_done=False,
  115. run_key=f"singleton-acquisition:{batch.id}:{run_key_suffix}",
  116. )
  117. with transaction(db_config) as conn:
  118. acquisition_repo = PostgresAcquisitionRepository(conn)
  119. decode_repo = PostgresDecodeRepository(conn)
  120. decode_service = DecodeService(settings=settings, repository=decode_repo)
  121. decode = run_decode_stage(
  122. candidate_repo=acquisition_repo,
  123. decode_service=decode_service,
  124. run_id=acquisition.run_id,
  125. limit=args.decode_limit,
  126. )
  127. ingest_records = _dry_ingest_payloads(decode_repo, decode.outputs)
  128. decoded_items = [str(output.item_id) for output in decode.outputs]
  129. first_item_id = decoded_items[0] if decoded_items else None
  130. summary = {
  131. "batch_id": str(batch.id),
  132. "run_id": str(acquisition.run_id),
  133. "active_family_keys": ["f1", "f2"],
  134. "query_count": query_count,
  135. "kept_count": kept_count,
  136. "persisted_queries": persisted_count,
  137. "platforms": list(platforms),
  138. "acquisition": _model_dump(acquisition),
  139. "decode": {
  140. "total": decode.total,
  141. "decoded": decode.decoded,
  142. "skipped": decode.skipped,
  143. "failed": decode.failed,
  144. "decoded_items": decoded_items,
  145. "payload_count": sum(len(output.payload_drafts) for output in decode.outputs),
  146. },
  147. "dry_run_ingest_count": len(ingest_records),
  148. "detail_url": (
  149. f"{args.frontend_base.rstrip('/')}/#/decode-item/{first_item_id}"
  150. if first_item_id
  151. else None
  152. ),
  153. }
  154. print(json.dumps(summary, ensure_ascii=False, indent=2, default=str))
  155. return 0 if first_item_id else 2
  156. if __name__ == "__main__":
  157. raise SystemExit(main())