test_acquisition_runner.py 9.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296
  1. from __future__ import annotations
  2. from uuid import UUID, uuid4
  3. from acquisition.classification.coarse import ClassificationResult
  4. from acquisition.domain import (
  5. AcquisitionJob,
  6. AcquisitionRun,
  7. CandidateItem,
  8. ItemClassification,
  9. MediaAsset,
  10. Query,
  11. QueryBatch,
  12. )
  13. from acquisition.media.service import StabilizedMedia
  14. from acquisition.platforms.base import PlatformCandidate, PlatformItem
  15. from acquisition.runner import run_batch
  16. from core.config import PgConfig, Settings
  17. def _settings() -> Settings:
  18. return Settings(
  19. pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
  20. aiddit_crawler_base_url="http://crawler.test",
  21. crawler_timeout=30,
  22. openrouter_timeout_seconds=90,
  23. openrouter_model="m",
  24. openrouter_base_url="http://openrouter.test",
  25. openrouter_api_key="k",
  26. llm_model="m",
  27. max_cards=12,
  28. frames_dir="f",
  29. douyin_ratio="540p",
  30. data_dir="data",
  31. )
  32. class FakeRepo:
  33. def __init__(self, *, job_status: str = "pending") -> None:
  34. self.batch_id = uuid4()
  35. self.query = Query(
  36. id=uuid4(),
  37. batch_id=self.batch_id,
  38. query_text="脚本 开头 怎么做",
  39. keep=True,
  40. status="ready",
  41. )
  42. self.job_status = job_status
  43. self.runs: list[AcquisitionRun] = []
  44. self.jobs: list[AcquisitionJob] = []
  45. self.updated_jobs: list[AcquisitionJob] = []
  46. self.items: list[CandidateItem] = []
  47. self.media: list[MediaAsset] = []
  48. self.classifications: list[ItemClassification] = []
  49. def create_query_batch(self, **kwargs):
  50. return QueryBatch(id=self.batch_id, **kwargs)
  51. def add_query(self, **kwargs):
  52. return Query(id=uuid4(), **kwargs)
  53. def list_queries_for_batch(self, batch_id: UUID, *, keep: bool | None = None):
  54. assert batch_id == self.batch_id
  55. assert keep is True
  56. return [self.query]
  57. def create_acquisition_run(self, **kwargs):
  58. run = AcquisitionRun(id=uuid4(), **kwargs)
  59. self.runs.append(run)
  60. return run
  61. def ensure_acquisition_job(self, **kwargs):
  62. kwargs["status"] = self.job_status
  63. job = AcquisitionJob(id=uuid4(), **kwargs)
  64. self.jobs.append(job)
  65. return job
  66. def update_acquisition_job(self, job_id: UUID, **kwargs):
  67. original = next(job for job in self.jobs if job.id == job_id)
  68. updated = original.model_copy(update=kwargs)
  69. self.updated_jobs.append(updated)
  70. return updated
  71. def upsert_candidate_item(self, **kwargs):
  72. item = CandidateItem(id=uuid4(), **kwargs)
  73. self.items.append(item)
  74. return item
  75. def add_media_asset(self, **kwargs):
  76. row = MediaAsset(id=uuid4(), **kwargs)
  77. self.media.append(row)
  78. return row
  79. def add_item_classification(self, **kwargs):
  80. row = ItemClassification(id=uuid4(), **kwargs)
  81. self.classifications.append(row)
  82. return row
  83. def get_run_summary(self, run_id: UUID):
  84. return {}
  85. def get_query_detail(self, *, run_id: UUID, query_id: UUID):
  86. return {}
  87. class FakeAdapter:
  88. platform = "weixin"
  89. def __init__(self) -> None:
  90. self.search_calls = 0
  91. self.detail_calls = 0
  92. def search(self, query, *, settings, limit, rate_limiter):
  93. self.search_calls += 1
  94. assert query == "脚本 开头 怎么做"
  95. assert limit == 2
  96. return [
  97. PlatformCandidate(
  98. rank=1,
  99. platform=self.platform,
  100. source_id="wx1",
  101. url="https://mp.weixin.qq.com/s/a",
  102. title="公众号脚本",
  103. author="作者",
  104. cover_url="https://img.test/a.jpg",
  105. )
  106. ]
  107. def fetch_detail(self, candidate, *, settings, rate_limiter):
  108. self.detail_calls += 1
  109. return PlatformItem(
  110. platform=self.platform,
  111. source_id=candidate.source_id,
  112. url=candidate.url,
  113. content_type="图文",
  114. title="公众号脚本",
  115. author="作者",
  116. body_text="先定受众再写开头",
  117. image_urls=["https://img.test/a.jpg"],
  118. raw={"ok": True},
  119. )
  120. class PartialAdapter(FakeAdapter):
  121. def search(self, query, *, settings, limit, rate_limiter):
  122. self.search_calls += 1
  123. return [
  124. PlatformCandidate(rank=1, platform=self.platform, source_id="bad", url="https://bad"),
  125. PlatformCandidate(rank=2, platform=self.platform, source_id="good", url="https://good"),
  126. ]
  127. def fetch_detail(self, candidate, *, settings, rate_limiter):
  128. self.detail_calls += 1
  129. if candidate.source_id == "bad":
  130. raise RuntimeError("detail boom")
  131. return PlatformItem(
  132. platform=self.platform,
  133. source_id=candidate.source_id,
  134. url=candidate.url,
  135. content_type="图文",
  136. title="可用详情",
  137. author="作者",
  138. body_text="先定受众",
  139. image_urls=[],
  140. raw={},
  141. )
  142. class FailingSearchAdapter(FakeAdapter):
  143. def search(self, query, *, settings, limit, rate_limiter):
  144. self.search_calls += 1
  145. raise RuntimeError("search boom")
  146. def test_run_batch_writes_jobs_items_media_and_classification():
  147. repo = FakeRepo()
  148. adapter = FakeAdapter()
  149. def media_stabilizer(**kwargs):
  150. assert kwargs["image_urls"] == ["https://img.test/a.jpg"]
  151. return [
  152. StabilizedMedia(
  153. media_type="image",
  154. source_url="https://img.test/a.jpg",
  155. cdn_url="https://cdn.test/a.jpg",
  156. position=1,
  157. )
  158. ]
  159. def classifier(**kwargs):
  160. assert kwargs["image_urls"] == ["https://cdn.test/a.jpg"]
  161. return ClassificationResult(
  162. is_creation_knowledge=True,
  163. label="creation",
  164. confidence=1.0,
  165. reason="是创作知识",
  166. knowledge="先定受众",
  167. prompt_version="pv1",
  168. result_payload={"knowledge": "先定受众"},
  169. )
  170. result = run_batch(
  171. repo,
  172. batch_id=repo.batch_id,
  173. settings=_settings(),
  174. platforms=("weixin",),
  175. search_limit=2,
  176. display_limit=1,
  177. adapter_factory=lambda platform: adapter,
  178. media_stabilizer=media_stabilizer,
  179. classifier=classifier,
  180. rate_limiter_factory=lambda platform: object(),
  181. )
  182. assert result.done == 1
  183. assert result.failed == 0
  184. assert result.total_jobs == 1
  185. assert adapter.search_calls == 1
  186. assert adapter.detail_calls == 1
  187. assert repo.jobs[0].query_id == repo.query.id
  188. assert repo.items[0].platform_item_id == "wx1"
  189. assert repo.media[0].cdn_url == "https://cdn.test/a.jpg"
  190. assert repo.classifications[0].is_creation_knowledge is True
  191. assert repo.updated_jobs[-1].status == "done"
  192. def test_run_batch_skip_done_does_not_call_platform():
  193. repo = FakeRepo(job_status="done")
  194. adapter = FakeAdapter()
  195. result = run_batch(
  196. repo,
  197. batch_id=repo.batch_id,
  198. settings=_settings(),
  199. platforms=("weixin",),
  200. adapter_factory=lambda platform: adapter,
  201. )
  202. assert result.skipped == 1
  203. assert adapter.search_calls == 0
  204. assert repo.updated_jobs == []
  205. def test_run_batch_marks_partial_when_some_details_fail():
  206. repo = FakeRepo()
  207. adapter = PartialAdapter()
  208. result = run_batch(
  209. repo,
  210. batch_id=repo.batch_id,
  211. settings=_settings(),
  212. platforms=("weixin",),
  213. search_limit=2,
  214. display_limit=2,
  215. adapter_factory=lambda platform: adapter,
  216. media_stabilizer=lambda **kwargs: [],
  217. classifier=lambda **kwargs: ClassificationResult(
  218. is_creation_knowledge=True,
  219. label="creation",
  220. confidence=1.0,
  221. reason="ok",
  222. ),
  223. rate_limiter_factory=lambda platform: object(),
  224. )
  225. assert result.partial == 1
  226. assert result.done == 0
  227. assert result.failed == 0
  228. assert repo.updated_jobs[-1].status == "partial"
  229. assert repo.updated_jobs[-1].metadata["display_count"] == 1
  230. assert repo.updated_jobs[-1].metadata["errors"] == ["detail boom"]
  231. def test_run_batch_marks_failed_when_search_fails_before_any_item():
  232. repo = FakeRepo()
  233. adapter = FailingSearchAdapter()
  234. result = run_batch(
  235. repo,
  236. batch_id=repo.batch_id,
  237. settings=_settings(),
  238. platforms=("weixin",),
  239. adapter_factory=lambda platform: adapter,
  240. rate_limiter_factory=lambda platform: object(),
  241. )
  242. assert result.failed == 1
  243. assert result.done == 0
  244. assert repo.updated_jobs[-1].status == "failed"
  245. assert "search boom" in (repo.updated_jobs[-1].error_message or "")
  246. def test_formal_runner_does_not_import_legacy_store():
  247. source = __import__("pathlib").Path("acquisition/runner.py").read_text(encoding="utf-8")
  248. assert "acquisition.store" not in source
  249. assert "creation_demo.json" not in source