| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296 |
- from __future__ import annotations
- from uuid import UUID, uuid4
- from acquisition.classification.coarse import ClassificationResult
- from acquisition.domain import (
- AcquisitionJob,
- AcquisitionRun,
- CandidateItem,
- ItemClassification,
- MediaAsset,
- Query,
- QueryBatch,
- )
- from acquisition.media.service import StabilizedMedia
- from acquisition.platforms.base import PlatformCandidate, PlatformItem
- from acquisition.runner import run_batch
- from core.config import PgConfig, Settings
- def _settings() -> Settings:
- return Settings(
- pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
- aiddit_crawler_base_url="http://crawler.test",
- crawler_timeout=30,
- openrouter_timeout_seconds=90,
- openrouter_model="m",
- openrouter_base_url="http://openrouter.test",
- openrouter_api_key="k",
- llm_model="m",
- max_cards=12,
- frames_dir="f",
- douyin_ratio="540p",
- data_dir="data",
- )
- class FakeRepo:
- def __init__(self, *, job_status: str = "pending") -> None:
- self.batch_id = uuid4()
- self.query = Query(
- id=uuid4(),
- batch_id=self.batch_id,
- query_text="脚本 开头 怎么做",
- keep=True,
- status="ready",
- )
- self.job_status = job_status
- self.runs: list[AcquisitionRun] = []
- self.jobs: list[AcquisitionJob] = []
- self.updated_jobs: list[AcquisitionJob] = []
- self.items: list[CandidateItem] = []
- self.media: list[MediaAsset] = []
- self.classifications: list[ItemClassification] = []
- def create_query_batch(self, **kwargs):
- return QueryBatch(id=self.batch_id, **kwargs)
- def add_query(self, **kwargs):
- return Query(id=uuid4(), **kwargs)
- def list_queries_for_batch(self, batch_id: UUID, *, keep: bool | None = None):
- assert batch_id == self.batch_id
- assert keep is True
- return [self.query]
- def create_acquisition_run(self, **kwargs):
- run = AcquisitionRun(id=uuid4(), **kwargs)
- self.runs.append(run)
- return run
- def ensure_acquisition_job(self, **kwargs):
- kwargs["status"] = self.job_status
- job = AcquisitionJob(id=uuid4(), **kwargs)
- self.jobs.append(job)
- return job
- def update_acquisition_job(self, job_id: UUID, **kwargs):
- original = next(job for job in self.jobs if job.id == job_id)
- updated = original.model_copy(update=kwargs)
- self.updated_jobs.append(updated)
- return updated
- def upsert_candidate_item(self, **kwargs):
- item = CandidateItem(id=uuid4(), **kwargs)
- self.items.append(item)
- return item
- def add_media_asset(self, **kwargs):
- row = MediaAsset(id=uuid4(), **kwargs)
- self.media.append(row)
- return row
- def add_item_classification(self, **kwargs):
- row = ItemClassification(id=uuid4(), **kwargs)
- self.classifications.append(row)
- return row
- def get_run_summary(self, run_id: UUID):
- return {}
- def get_query_detail(self, *, run_id: UUID, query_id: UUID):
- return {}
- class FakeAdapter:
- platform = "weixin"
- def __init__(self) -> None:
- self.search_calls = 0
- self.detail_calls = 0
- def search(self, query, *, settings, limit, rate_limiter):
- self.search_calls += 1
- assert query == "脚本 开头 怎么做"
- assert limit == 2
- return [
- PlatformCandidate(
- rank=1,
- platform=self.platform,
- source_id="wx1",
- url="https://mp.weixin.qq.com/s/a",
- title="公众号脚本",
- author="作者",
- cover_url="https://img.test/a.jpg",
- )
- ]
- def fetch_detail(self, candidate, *, settings, rate_limiter):
- self.detail_calls += 1
- return PlatformItem(
- platform=self.platform,
- source_id=candidate.source_id,
- url=candidate.url,
- content_type="图文",
- title="公众号脚本",
- author="作者",
- body_text="先定受众再写开头",
- image_urls=["https://img.test/a.jpg"],
- raw={"ok": True},
- )
- class PartialAdapter(FakeAdapter):
- def search(self, query, *, settings, limit, rate_limiter):
- self.search_calls += 1
- return [
- PlatformCandidate(rank=1, platform=self.platform, source_id="bad", url="https://bad"),
- PlatformCandidate(rank=2, platform=self.platform, source_id="good", url="https://good"),
- ]
- def fetch_detail(self, candidate, *, settings, rate_limiter):
- self.detail_calls += 1
- if candidate.source_id == "bad":
- raise RuntimeError("detail boom")
- return PlatformItem(
- platform=self.platform,
- source_id=candidate.source_id,
- url=candidate.url,
- content_type="图文",
- title="可用详情",
- author="作者",
- body_text="先定受众",
- image_urls=[],
- raw={},
- )
- class FailingSearchAdapter(FakeAdapter):
- def search(self, query, *, settings, limit, rate_limiter):
- self.search_calls += 1
- raise RuntimeError("search boom")
- def test_run_batch_writes_jobs_items_media_and_classification():
- repo = FakeRepo()
- adapter = FakeAdapter()
- def media_stabilizer(**kwargs):
- assert kwargs["image_urls"] == ["https://img.test/a.jpg"]
- return [
- StabilizedMedia(
- media_type="image",
- source_url="https://img.test/a.jpg",
- cdn_url="https://cdn.test/a.jpg",
- position=1,
- )
- ]
- def classifier(**kwargs):
- assert kwargs["image_urls"] == ["https://cdn.test/a.jpg"]
- return ClassificationResult(
- is_creation_knowledge=True,
- label="creation",
- confidence=1.0,
- reason="是创作知识",
- knowledge="先定受众",
- prompt_version="pv1",
- result_payload={"knowledge": "先定受众"},
- )
- result = run_batch(
- repo,
- batch_id=repo.batch_id,
- settings=_settings(),
- platforms=("weixin",),
- search_limit=2,
- display_limit=1,
- adapter_factory=lambda platform: adapter,
- media_stabilizer=media_stabilizer,
- classifier=classifier,
- rate_limiter_factory=lambda platform: object(),
- )
- assert result.done == 1
- assert result.failed == 0
- assert result.total_jobs == 1
- assert adapter.search_calls == 1
- assert adapter.detail_calls == 1
- assert repo.jobs[0].query_id == repo.query.id
- assert repo.items[0].platform_item_id == "wx1"
- assert repo.media[0].cdn_url == "https://cdn.test/a.jpg"
- assert repo.classifications[0].is_creation_knowledge is True
- assert repo.updated_jobs[-1].status == "done"
- def test_run_batch_skip_done_does_not_call_platform():
- repo = FakeRepo(job_status="done")
- adapter = FakeAdapter()
- result = run_batch(
- repo,
- batch_id=repo.batch_id,
- settings=_settings(),
- platforms=("weixin",),
- adapter_factory=lambda platform: adapter,
- )
- assert result.skipped == 1
- assert adapter.search_calls == 0
- assert repo.updated_jobs == []
- def test_run_batch_marks_partial_when_some_details_fail():
- repo = FakeRepo()
- adapter = PartialAdapter()
- result = run_batch(
- repo,
- batch_id=repo.batch_id,
- settings=_settings(),
- platforms=("weixin",),
- search_limit=2,
- display_limit=2,
- adapter_factory=lambda platform: adapter,
- media_stabilizer=lambda **kwargs: [],
- classifier=lambda **kwargs: ClassificationResult(
- is_creation_knowledge=True,
- label="creation",
- confidence=1.0,
- reason="ok",
- ),
- rate_limiter_factory=lambda platform: object(),
- )
- assert result.partial == 1
- assert result.done == 0
- assert result.failed == 0
- assert repo.updated_jobs[-1].status == "partial"
- assert repo.updated_jobs[-1].metadata["display_count"] == 1
- assert repo.updated_jobs[-1].metadata["errors"] == ["detail boom"]
- def test_run_batch_marks_failed_when_search_fails_before_any_item():
- repo = FakeRepo()
- adapter = FailingSearchAdapter()
- result = run_batch(
- repo,
- batch_id=repo.batch_id,
- settings=_settings(),
- platforms=("weixin",),
- adapter_factory=lambda platform: adapter,
- rate_limiter_factory=lambda platform: object(),
- )
- assert result.failed == 1
- assert result.done == 0
- assert repo.updated_jobs[-1].status == "failed"
- assert "search boom" in (repo.updated_jobs[-1].error_message or "")
- def test_formal_runner_does_not_import_legacy_store():
- source = __import__("pathlib").Path("acquisition/runner.py").read_text(encoding="utf-8")
- assert "acquisition.store" not in source
- assert "creation_demo.json" not in source
|