| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692 |
- from __future__ import annotations
- import asyncio
- import importlib
- import json
- from contextlib import contextmanager
- from unittest.mock import patch
- import pytest
- from sqlalchemy import create_engine, event, select
- from sqlalchemy.orm import Session, sessionmaker
- from agents.find_agent import create_find_agent
- from agents.find_agent.async_runner import arun_find_agent
- from agents.find_agent.demand_run import (
- FindDemandContext,
- FindDemandPoint,
- FindDemandVideo,
- build_find_agent_user_input,
- prepare_video_discovery_run,
- )
- from agents.find_agent.tools import video_discovery_store
- from agents.find_agent.tools.batch_search_and_record import batch_search_and_record
- from supply_infra.db.models.video_discovery import (
- VideoDiscoveryCandidate,
- VideoDiscoveryRun,
- VideoDiscoverySearch,
- )
- from supply_infra.db.repositories.video_discovery_repo import (
- VideoDiscoveryRepository,
- )
- from supply_infra.scheduler.jobs.discover_videos_from_demands import (
- discover_videos_from_demands,
- )
- def _expire_on_commit_session_factory() -> sessionmaker[Session]:
- """Match production get_session: commit expires ORM instances."""
- engine = create_engine("sqlite+pysqlite:///:memory:")
- # SQLite 对 BigInteger PK 不会自增,测试里显式写入 id。
- VideoDiscoveryRun.__table__.create(engine)
- VideoDiscoveryCandidate.__table__.create(engine)
- return sessionmaker(bind=engine, autoflush=False, autocommit=False)
- def _patch_service_session(
- monkeypatch: pytest.MonkeyPatch,
- factory: sessionmaker[Session],
- ) -> None:
- @contextmanager
- def get_test_session():
- session = factory()
- try:
- yield session
- session.commit()
- except Exception:
- session.rollback()
- raise
- finally:
- session.close()
- monkeypatch.setattr(
- "supply_infra.services.video_discovery_service.get_session",
- get_test_session,
- )
- import supply_infra.services.video_discovery_service as service_module
- service_module._default_service = None
- def _seed_run(
- factory: sessionmaker[Session],
- *,
- run_id: str,
- demand_grade_id: int,
- status: str = "running",
- row_id: int = 1,
- ) -> None:
- with factory() as session:
- session.add(
- VideoDiscoveryRun(
- id=row_id,
- run_id=run_id,
- biz_dt="20260728",
- demand_grade_id=demand_grade_id,
- demand_word="广场舞",
- seed_video_id="vid-1",
- seed_video_title="参考标题",
- relevant_points_json="[]",
- status=status,
- )
- )
- session.commit()
- def test_create_video_discovery_run_reuses_precreated_run(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- """定时任务预创建 run 后,Agent 复用 run_id 不得触发 DetachedInstanceError。"""
- factory = _expire_on_commit_session_factory()
- _patch_service_session(monkeypatch, factory)
- _seed_run(factory, run_id="precreated-run", demand_grade_id=101)
- payload = json.loads(
- video_discovery_store.create_video_discovery_run(
- demand_word="广场舞",
- seed_video_title="参考标题",
- relevant_points=[],
- run_id="precreated-run",
- demand_grade_id=101,
- )
- )
- assert "error" not in payload
- assert payload["run_id"] == "precreated-run"
- assert payload["status"] == "running"
- assert payload["pre_created"] is True
- def test_prepare_then_reuse_run_id_scheduled_flow(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- """完整调度衔接:prepare 得到 run_id → create 复用,全程无 DetachedInstanceError。"""
- factory = _expire_on_commit_session_factory()
- _patch_service_session(monkeypatch, factory)
- # 先写入一条 failed 记录,prepare 会原地重置为 running 并复用 run_id。
- _seed_run(
- factory,
- run_id="scheduled-run",
- demand_grade_id=202,
- status="failed",
- )
- ctx = FindDemandContext(
- biz_dt="20260728",
- demand_grade_id=202,
- demand_name="广场舞",
- grade="S",
- videos=[
- FindDemandVideo(
- video_id="vid-1",
- title="参考标题",
- points=[
- FindDemandPoint(
- point="动作简单",
- point_type="key",
- )
- ],
- )
- ],
- )
- run_id, skip_reason = prepare_video_discovery_run(ctx)
- assert skip_reason is None
- assert run_id == "scheduled-run"
- reuse_id, reuse_skip = prepare_video_discovery_run(ctx)
- assert reuse_id == "scheduled-run"
- assert reuse_skip is None
- payload = json.loads(
- video_discovery_store.create_video_discovery_run(
- demand_word=ctx.demand_name,
- seed_video_title="参考标题",
- relevant_points=[{"point": "动作简单", "point_type": "key"}],
- run_id=run_id,
- demand_grade_id=ctx.demand_grade_id,
- seed_video_id="vid-1",
- )
- )
- assert "error" not in payload
- assert payload["pre_created"] is True
- assert payload["run_id"] == run_id
- def test_prepare_skips_when_run_already_finished(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- factory = _expire_on_commit_session_factory()
- _patch_service_session(monkeypatch, factory)
- _seed_run(
- factory,
- run_id="finished-run",
- demand_grade_id=303,
- status="finished",
- )
- ctx = FindDemandContext(
- biz_dt="20260728",
- demand_grade_id=303,
- demand_name="广场舞",
- grade="S",
- videos=[
- FindDemandVideo(
- video_id="vid-1",
- title="参考标题",
- points=[FindDemandPoint(point="动作简单", point_type="key")],
- )
- ],
- )
- run_id, skip_reason = prepare_video_discovery_run(ctx)
- assert run_id is None
- assert skip_reason is not None
- assert "finished" in skip_reason
- def test_video_discovery_models_exclude_unused_columns() -> None:
- run_columns = set(VideoDiscoveryRun.__table__.columns.keys())
- candidate_columns = set(VideoDiscoveryCandidate.__table__.columns.keys())
- assert "backup_count" not in run_columns
- assert {
- "video_url",
- "content_analysis",
- "content_analysis_verified",
- "hit_points_json",
- "publish_timestamp",
- "detail_verified",
- "content_portrait_attempted",
- "account_portrait_attempted",
- "age_portraits_normalized",
- "expansion_worthy_tags_json",
- "confidence",
- "relevance_reason",
- "elder_reason",
- "share_reason",
- "manual_review_note",
- "manual_review_status",
- }.isdisjoint(candidate_columns)
- assert "search_id" in candidate_columns
- candidate_constraints = {
- constraint.name
- for constraint in VideoDiscoveryCandidate.__table__.constraints
- }
- search_constraints = {
- constraint.name
- for constraint in VideoDiscoverySearch.__table__.constraints
- }
- assert "uk_video_discovery_candidate_run_aweme" not in candidate_constraints
- assert "fk_video_discovery_candidate_search" in candidate_constraints
- assert "uk_video_discovery_search_key" not in search_constraints
- def test_create_find_agent_registers_discovery_tools() -> None:
- agent = create_find_agent()
- assert agent.name == "find_agent"
- assert "batch_search_and_record" in agent.tools.list_tools()
- assert "batch_update_video_discovery_candidates" in agent.tools.list_tools()
- assert "update_video_discovery_run_status" in agent.tools.list_tools()
- assert "create_video_discovery_run" not in agent.tools.list_tools()
- assert "batch_save_video_candidate_evaluations" not in agent.tools.list_tools()
- assert "audit_video_discovery_run" not in agent.tools.list_tools()
- assert "query_video_discovery_state" in agent.tools.list_tools()
- def test_find_agent_input_uses_reference_videos_without_seed_fields() -> None:
- ctx = FindDemandContext(
- biz_dt="20260729",
- demand_grade_id=101,
- demand_name="广场舞",
- grade="S",
- videos=[
- FindDemandVideo(
- video_id="vid-1",
- title="参考标题",
- points=[FindDemandPoint(point="动作简单", point_type="key")],
- )
- ],
- )
- user_input = build_find_agent_user_input(ctx, "scheduled-run")
- assert "seed_video_id:" not in user_input
- assert "seed_video_title:" not in user_input
- assert "reference_videos:" in user_input
- assert '"video_id": "vid-1"' in user_input
- assert '"title": "参考标题"' in user_input
- assert "create_video_discovery_run" not in user_input
- assert "relevant_points" not in user_input
- @pytest.mark.asyncio
- async def test_douyin_search_automatically_persists_page(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- search_module = importlib.import_module(
- "agents.find_agent.tools.douyin_search"
- )
- async def fake_raw_search(**_kwargs):
- return json.dumps(
- {
- "results_count": 1,
- "has_more": False,
- "search_results": [{"aweme_id": "auto-saved"}],
- }
- )
- persisted: dict[str, object] = {}
- def fake_persist(payload_json: str, **kwargs):
- persisted.update(kwargs)
- payload = json.loads(payload_json)
- payload.update(
- {
- "persisted": True,
- "search_id": 11,
- "new_candidate_count": 1,
- "candidates": [
- {
- "candidate_id": 21,
- "search_id": 11,
- "aweme_id": "auto-saved",
- "title": "自动保存",
- "decision_bucket": "pending_evaluation",
- }
- ],
- }
- )
- return json.dumps(payload)
- monkeypatch.setattr(search_module, "_douyin_search_raw", fake_raw_search)
- monkeypatch.setattr(search_module, "persist_search_payload", fake_persist)
- result = json.loads(
- await search_module.douyin_search(
- run_id="run-auto-save",
- keyword="广场舞",
- query_reason="验证需求根搜索",
- source_type="demand",
- )
- )
- assert result["persisted"] is True
- assert result["search_id"] == 11
- assert result["candidates"][0]["candidate_id"] == 21
- assert persisted["run_id"] == "run-auto-save"
- assert persisted["keyword"] == "广场舞"
- assert persisted["provider"] == "internal_keyword"
- @pytest.mark.asyncio
- async def test_batch_search_records_each_page_and_carries_parent(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- batch_module = importlib.import_module(
- "agents.find_agent.tools.batch_search_and_record"
- )
- calls: list[dict[str, object]] = []
- async def fake_search(**kwargs):
- calls.append(kwargs)
- page_no = int(kwargs["page_no"])
- return json.dumps(
- {
- "results_count": 1,
- "has_more": page_no == 1,
- "next_cursor": "next-page" if page_no == 1 else None,
- "persisted": True,
- "search_id": 100 + page_no,
- "new_candidate_count": 1,
- "candidates": [
- {
- "candidate_id": 200 + page_no,
- "search_id": 100 + page_no,
- "aweme_id": "same-video",
- "title": f"第 {page_no} 页",
- "decision_bucket": "pending_evaluation",
- }
- ],
- }
- )
- monkeypatch.setattr(batch_module, "douyin_search", fake_search)
- result = json.loads(
- await batch_search_and_record(
- run_id="run-batch",
- searches=[
- {
- "keyword": "广场舞",
- "query_reason": "验证需求根搜索",
- "source_type": "demand",
- "max_pages": 2,
- }
- ],
- )
- )
- assert result["saved_page_count"] == 2
- assert result["new_candidate_count"] == 2
- assert result["tasks"][0]["pages"][0]["candidates"][0]["candidate_id"] == 201
- assert result["tasks"][0]["pages"][1]["candidates"][0]["candidate_id"] == 202
- assert calls[0]["parent_search_id"] is None
- assert calls[1]["parent_search_id"] == 101
- assert calls[1]["cursor"] == "next-page"
- def test_batch_update_candidates_uses_database_candidate_id(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- captured: dict[str, object] = {}
- class FakeService:
- def update_candidates(self, run_id, rows):
- captured["run_id"] = run_id
- captured["rows"] = rows
- return {
- "updated_count": 1,
- "candidates": [
- {
- "candidate_id": 901,
- "search_id": 801,
- "aweme_id": "same-video",
- "decision_bucket": "primary",
- }
- ],
- }
- monkeypatch.setattr(
- video_discovery_store,
- "get_video_discovery_service",
- lambda: FakeService(),
- )
- result = json.loads(
- video_discovery_store.batch_update_video_discovery_candidates(
- run_id="run-update",
- items=[
- {
- "candidate_id": 901,
- "decision_bucket": "primary",
- "relevance_score": 0.8,
- "elder_score": 0.7,
- "share_score": 0.6,
- }
- ],
- )
- )
- assert result["updated_count"] == 1
- assert captured["run_id"] == "run-update"
- assert captured["rows"][0]["candidate_id"] == 901
- assert "aweme_id" not in captured["rows"][0]
- def test_each_search_inserts_new_candidate_occurrences() -> None:
- engine = create_engine("sqlite+pysqlite:///:memory:")
- VideoDiscoveryRun.__table__.create(engine)
- VideoDiscoverySearch.__table__.create(engine)
- VideoDiscoveryCandidate.__table__.create(engine)
- factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
- ids = {"search": 100, "candidate": 1000}
- @event.listens_for(factory.class_, "before_flush")
- def assign_sqlite_bigint_ids(session, _flush_context, _instances):
- for entity in session.new:
- if isinstance(entity, VideoDiscoverySearch) and entity.id is None:
- ids["search"] += 1
- entity.id = ids["search"]
- elif isinstance(entity, VideoDiscoveryCandidate) and entity.id is None:
- ids["candidate"] += 1
- entity.id = ids["candidate"]
- with factory() as session:
- session.add(
- VideoDiscoveryRun(
- id=1,
- run_id="run-occurrences",
- demand_word="广场舞",
- relevant_points_json="[]",
- status="running",
- )
- )
- session.commit()
- search_values = {
- "run_id": "run-occurrences",
- "search_key": "same-search-key",
- "keyword": "广场舞",
- "query_reason": "验证相同搜索也生成新记录",
- "source_type": "demand",
- "provider": "internal_keyword",
- "content_type": "视频",
- "sort_type": "综合排序",
- "publish_time": "不限",
- "cursor": "0",
- "page_no": 1,
- "results_count": 1,
- "new_candidate_count": 0,
- "has_more": 0,
- "status": "success",
- }
- candidate_rows = [
- {
- "aweme_id": "same-video",
- "title": "同一个视频",
- "_source_keyword": "广场舞",
- }
- ]
- with factory() as session:
- repo = VideoDiscoveryRepository(session)
- first_search, first_candidates = repo.save_search_page(
- dict(search_values),
- candidate_rows,
- )
- second_search, second_candidates = repo.save_search_page(
- dict(search_values),
- candidate_rows,
- )
- session.commit()
- assert first_search.id != second_search.id
- assert first_candidates[0].id != second_candidates[0].id
- assert first_candidates[0].search_id == first_search.id
- assert second_candidates[0].search_id == second_search.id
- with factory() as session:
- searches = session.scalars(select(VideoDiscoverySearch)).all()
- candidates = session.scalars(select(VideoDiscoveryCandidate)).all()
- assert len(searches) == 2
- assert len(candidates) == 2
- assert {candidate.aweme_id for candidate in candidates} == {"same-video"}
- def _seed_candidate(
- factory: sessionmaker[Session],
- *,
- run_id: str,
- aweme_id: str = "7631830155522179258",
- row_id: int = 1,
- ) -> None:
- with factory() as session:
- session.add(
- VideoDiscoveryCandidate(
- id=row_id,
- run_id=run_id,
- aweme_id=aweme_id,
- decision_bucket="pending_evaluation",
- )
- )
- session.commit()
- def test_list_skip_grade_ids_when_candidates_exist(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- factory = _expire_on_commit_session_factory()
- _patch_service_session(monkeypatch, factory)
- _seed_run(factory, run_id="running-empty", demand_grade_id=301, status="running")
- _seed_run(
- factory,
- run_id="running-with-candidates",
- demand_grade_id=302,
- status="running",
- row_id=2,
- )
- _seed_candidate(factory, run_id="running-with-candidates")
- from supply_infra.services.video_discovery_service import get_video_discovery_service
- skip_ids = get_video_discovery_service().list_skip_grade_ids("20260728")
- assert skip_ids == {302}
- def test_evaluate_find_agent_run_succeeds_when_candidates_exist(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- from agents.find_agent.run_outcome import evaluate_find_agent_run
- from supply_agent.types import AgentResult
- monkeypatch.setattr(
- "agents.find_agent.run_outcome.get_video_discovery_service",
- lambda: type(
- "Svc",
- (),
- {"has_candidates": staticmethod(lambda _run_id: True)},
- )(),
- )
- outcome = evaluate_find_agent_run(
- "run-1",
- AgentResult(content="任意文案", messages=[], iterations=3, tool_calls_made=0),
- )
- assert outcome.succeeded is True
- assert outcome.failure_reason is None
- def test_evaluate_find_agent_run_fails_without_candidates(
- monkeypatch: pytest.MonkeyPatch,
- ) -> None:
- from agents.find_agent.run_outcome import evaluate_find_agent_run
- from supply_agent.types import AgentResult
- monkeypatch.setattr(
- "agents.find_agent.run_outcome.get_video_discovery_service",
- lambda: type(
- "Svc",
- (),
- {"has_candidates": staticmethod(lambda _run_id: False)},
- )(),
- )
- outcome = evaluate_find_agent_run(
- "run-1",
- AgentResult(
- content="任务未完成(工具故障)",
- messages=[],
- iterations=1,
- tool_calls_made=0,
- ),
- )
- assert outcome.succeeded is False
- assert outcome.failure_reason == "no_candidates"
- @patch(
- "supply_infra.scheduler.jobs.discover_videos_from_demands.process_single_discover"
- )
- @patch("supply_infra.scheduler.jobs.discover_videos_from_demands._count_passed_videos")
- @patch(
- "supply_infra.scheduler.jobs.discover_videos_from_demands.filter_pending_contexts"
- )
- @patch(
- "supply_infra.scheduler.jobs.discover_videos_from_demands.list_find_demand_contexts"
- )
- def test_stops_discovery_after_200_passed_videos(
- mock_list_contexts,
- mock_filter_contexts,
- mock_count_passed,
- mock_process,
- ) -> None:
- contexts = [
- FindDemandContext(
- biz_dt="20260727",
- demand_grade_id=1,
- demand_name="需求A",
- grade="S",
- ),
- FindDemandContext(
- biz_dt="20260727",
- demand_grade_id=2,
- demand_name="需求B",
- grade="A",
- ),
- ]
- mock_list_contexts.return_value = ("20260727", contexts)
- mock_filter_contexts.return_value = (
- contexts,
- {"total_loaded": 2, "skipped_already_done": 0},
- )
- mock_count_passed.side_effect = [199, 200]
- mock_process.return_value = {"success": True, "skipped": False}
- result = discover_videos_from_demands("20260727", workers=1)
- assert mock_process.call_count == 1
- assert result["processed"] == 1
- assert result["passed_videos"] == 200
- assert result["stopped_by_passed_video_limit"] is True
- class _AsyncClient:
- def __init__(self) -> None:
- self.closed = False
- async def close(self) -> None:
- self.closed = True
- class _SlowAgent:
- def __init__(self) -> None:
- self.llm = type("LLM", (), {"_async_client": _AsyncClient()})()
- async def arun_core(self, _user_input: str) -> None:
- await asyncio.sleep(60)
- @pytest.mark.asyncio
- async def test_find_agent_timeout_closes_async_client() -> None:
- agent = _SlowAgent()
- with pytest.raises(TimeoutError, match="find_agent timed out"):
- await arun_find_agent(
- agent, # type: ignore[arg-type]
- "test",
- run_id="test-run",
- timeout_seconds=0.01,
- )
- assert agent.llm._async_client.closed is True
|