runner.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321
  1. """Formal acquisition runner backed by the repository boundary."""
  2. from __future__ import annotations
  3. from dataclasses import dataclass
  4. from typing import Any, Callable
  5. from uuid import UUID
  6. from acquisition.classification.coarse import ClassificationResult, coarse_classify_item
  7. from acquisition.crawler import RateLimiter
  8. from acquisition.domain import AcquisitionJob, Query
  9. from acquisition.media.service import StabilizedMedia, stabilize_media_urls
  10. from acquisition.platforms import PlatformAdapter, get_platform_adapter
  11. from acquisition.repositories.base import AcquisitionRepository
  12. from core.config import Settings
  13. DEFAULT_PLATFORMS = ("xiaohongshu", "weixin", "douyin")
  14. AdapterFactory = Callable[[str], PlatformAdapter]
  15. Classifier = Callable[..., ClassificationResult]
  16. MediaStabilizer = Callable[..., list[StabilizedMedia]]
  17. RateLimiterFactory = Callable[[str], Any]
  18. @dataclass(frozen=True)
  19. class RunBatchResult:
  20. run_id: UUID
  21. total_jobs: int
  22. done: int
  23. partial: int
  24. failed: int
  25. skipped: int
  26. class PlatformRateLimiter:
  27. """Share one rate bucket per platform across search and detail calls."""
  28. def __init__(self, platform: str, delegate: RateLimiter | None = None) -> None:
  29. self.platform = platform
  30. self.delegate = delegate or RateLimiter(
  31. min_interval_seconds=10.0,
  32. max_interval_seconds=12.0,
  33. )
  34. def wait(self, bucket: str) -> None:
  35. self.delegate.wait(self.platform)
  36. def _query_id(query: Query) -> UUID:
  37. if query.id is None:
  38. raise RuntimeError("formal acquisition queries must have id before running")
  39. return query.id
  40. def _job_id(job: AcquisitionJob) -> UUID:
  41. if job.id is None:
  42. raise RuntimeError("formal acquisition jobs must have id before running")
  43. return job.id
  44. def _best_video_url(media: list[StabilizedMedia]) -> str:
  45. for row in media:
  46. if row.media_type == "video" and row.cdn_url:
  47. return row.cdn_url
  48. return ""
  49. def _image_urls(media: list[StabilizedMedia]) -> list[str]:
  50. return [row.cdn_url for row in media if row.media_type == "image" and row.cdn_url]
  51. def _source_payload(candidate: Any, detail: Any) -> dict[str, Any]:
  52. return {
  53. "candidate": candidate.model_dump() if hasattr(candidate, "model_dump") else {},
  54. "detail": detail.raw if isinstance(getattr(detail, "raw", None), dict) else {},
  55. }
  56. def _record_candidate(
  57. repo: AcquisitionRepository,
  58. *,
  59. job: AcquisitionJob,
  60. query: Query,
  61. platform: str,
  62. candidate: Any,
  63. detail: Any,
  64. settings: Settings,
  65. media_stabilizer: MediaStabilizer,
  66. classifier: Classifier,
  67. classify: bool,
  68. ) -> bool:
  69. media_rows = media_stabilizer(
  70. image_urls=detail.image_urls,
  71. video_urls=detail.video_urls,
  72. settings=settings,
  73. )
  74. item = repo.upsert_candidate_item(
  75. platform=platform,
  76. job_id=_job_id(job),
  77. query_id=_query_id(query),
  78. platform_item_id=detail.source_id or candidate.source_id or None,
  79. canonical_url=detail.url or candidate.url or None,
  80. content_type=detail.content_type or None,
  81. title=detail.title or None,
  82. author_name=detail.author or None,
  83. raw_summary=(detail.body_text or "")[:1000],
  84. status="candidate",
  85. source_payload=_source_payload(candidate, detail),
  86. metadata={"candidate_rank": candidate.rank},
  87. )
  88. if item.id is None:
  89. raise RuntimeError("repository returned candidate without id")
  90. for row in media_rows:
  91. repo.add_media_asset(
  92. item_id=item.id,
  93. media_type=row.media_type,
  94. source_url=row.source_url,
  95. oss_url=row.cdn_url,
  96. cdn_url=row.cdn_url,
  97. position=row.position,
  98. status=row.status,
  99. source_payload={},
  100. metadata={},
  101. )
  102. if classify:
  103. result = classifier(
  104. platform=platform,
  105. title=detail.title,
  106. body_text=detail.body_text,
  107. image_urls=_image_urls(media_rows),
  108. video_url=_best_video_url(media_rows),
  109. settings=settings,
  110. )
  111. repo.add_item_classification(
  112. item_id=item.id,
  113. is_creation_knowledge=result.is_creation_knowledge,
  114. label=result.label,
  115. confidence=result.confidence,
  116. reason=result.reason,
  117. model_name=settings.video_model,
  118. prompt_version=result.prompt_version,
  119. result_payload=result.result_payload,
  120. status=result.status,
  121. error_message=result.error_message,
  122. )
  123. return True
  124. def run_batch(
  125. repo: AcquisitionRepository,
  126. *,
  127. batch_id: UUID,
  128. settings: Settings,
  129. platforms: tuple[str, ...] | list[str] = DEFAULT_PLATFORMS,
  130. search_limit: int = 10,
  131. display_limit: int = 5,
  132. classify: bool = True,
  133. resume: bool = True,
  134. skip_done: bool = True,
  135. run_key: str | None = None,
  136. adapter_factory: AdapterFactory = get_platform_adapter,
  137. media_stabilizer: MediaStabilizer = stabilize_media_urls,
  138. classifier: Classifier = coarse_classify_item,
  139. rate_limiter_factory: RateLimiterFactory | None = None,
  140. ) -> RunBatchResult:
  141. """Run query x platform acquisition and write formal cloud-state rows."""
  142. queries = repo.list_queries_for_batch(batch_id, keep=True)
  143. run = repo.create_acquisition_run(
  144. batch_id=batch_id,
  145. run_key=run_key or f"acquisition:{batch_id}",
  146. status="running",
  147. metadata={
  148. "platforms": list(platforms),
  149. "search_limit": search_limit,
  150. "display_limit": display_limit,
  151. "classify": classify,
  152. "resume": resume,
  153. },
  154. )
  155. if run.id is None:
  156. raise RuntimeError("repository returned acquisition run without id")
  157. total_jobs = len(queries) * len(platforms)
  158. done = partial = failed = skipped = 0
  159. gates: dict[str, Any] = {}
  160. for query in queries:
  161. query_id = _query_id(query)
  162. for platform in platforms:
  163. job = repo.ensure_acquisition_job(
  164. run_id=run.id,
  165. query_id=query_id,
  166. platform=platform,
  167. search_limit=search_limit,
  168. display_limit=display_limit,
  169. status="pending",
  170. metadata={"query_text": query.query_text},
  171. )
  172. if skip_done and job.status == "done":
  173. skipped += 1
  174. continue
  175. attempts = job.attempt_count + 1
  176. job = repo.update_acquisition_job(
  177. _job_id(job),
  178. status="running",
  179. attempt_count=attempts,
  180. error_message=None,
  181. )
  182. errors: list[str] = []
  183. display_count = 0
  184. searched_count = 0
  185. try:
  186. adapter = adapter_factory(platform)
  187. gate = gates.get(platform)
  188. if gate is None:
  189. gate = (
  190. rate_limiter_factory(platform)
  191. if rate_limiter_factory
  192. else PlatformRateLimiter(platform)
  193. )
  194. gates[platform] = gate
  195. candidates = adapter.search(
  196. query.query_text,
  197. settings=settings,
  198. limit=search_limit,
  199. rate_limiter=gate,
  200. )
  201. searched_count = len(candidates)
  202. for candidate in candidates:
  203. if display_count >= display_limit:
  204. break
  205. try:
  206. detail = adapter.fetch_detail(
  207. candidate,
  208. settings=settings,
  209. rate_limiter=gate,
  210. )
  211. if _record_candidate(
  212. repo,
  213. job=job,
  214. query=query,
  215. platform=platform,
  216. candidate=candidate,
  217. detail=detail,
  218. settings=settings,
  219. media_stabilizer=media_stabilizer,
  220. classifier=classifier,
  221. classify=classify,
  222. ):
  223. display_count += 1
  224. except Exception as exc:
  225. errors.append(str(exc)[:160])
  226. continue
  227. status = "done" if display_count >= display_limit else (
  228. "partial" if display_count else "failed"
  229. )
  230. if status == "done":
  231. done += 1
  232. elif status == "partial":
  233. partial += 1
  234. else:
  235. failed += 1
  236. repo.update_acquisition_job(
  237. _job_id(job),
  238. status=status,
  239. attempt_count=attempts,
  240. error_message=None if status != "failed" else "; ".join(errors[-3:]),
  241. metadata={
  242. "query_text": query.query_text,
  243. "searched_count": searched_count,
  244. "display_count": display_count,
  245. "errors": errors[-3:],
  246. },
  247. )
  248. except Exception as exc:
  249. failed += 1
  250. repo.update_acquisition_job(
  251. _job_id(job),
  252. status="failed",
  253. attempt_count=attempts,
  254. error_message=str(exc)[:300],
  255. metadata={
  256. "query_text": query.query_text,
  257. "searched_count": searched_count,
  258. "display_count": display_count,
  259. "errors": errors[-3:],
  260. },
  261. )
  262. run_status = (
  263. "done"
  264. if failed == 0 and partial == 0
  265. else "partial"
  266. if done > 0 or partial > 0
  267. else "failed"
  268. )
  269. update_run = getattr(repo, "update_acquisition_run", None)
  270. if update_run:
  271. update_run(
  272. run.id,
  273. status=run_status,
  274. metadata={
  275. "total_jobs": total_jobs,
  276. "done": done,
  277. "partial": partial,
  278. "failed": failed,
  279. "skipped": skipped,
  280. },
  281. )
  282. return RunBatchResult(
  283. run_id=run.id,
  284. total_jobs=total_jobs,
  285. done=done,
  286. partial=partial,
  287. failed=failed,
  288. skipped=skipped,
  289. )