from __future__ import annotations import asyncio import importlib import json from contextlib import contextmanager from datetime import datetime from decimal import Decimal 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 agents.find_agent.support.search_persistence import _candidate_from_search_result from agents.find_agent.support.douyin_search import ( DEFAULT_MIN_DURATION_SECONDS as INTERNAL_SEARCH_MIN_DURATION, ) from agents.find_agent.support.douyin_search_tikhub import ( DEFAULT_MIN_DURATION_SECONDS as TIKHUB_SEARCH_MIN_DURATION, ) from agents.find_agent.support.douyin_user_videos import ( DEFAULT_MIN_DURATION_SECONDS as AUTHOR_SEARCH_MIN_DURATION, ) from supply_infra.video_discovery_gates import evaluate_candidate_gate 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 assert "current_datetime:" in user_input assert "timezone:Asia/Shanghai" in user_input assert "quality_gate_rules:" in user_input _P0_RULES = { "rule_version": "test-p0", "timezone": "Asia/Shanghai", "current_datetime": "2026-07-31T12:00:00+08:00", "current_date": "2026-07-31", "min_duration_seconds": 30, "min_share_count": 1000, "min_content_50_plus_ratio": 0.20, "min_account_50_plus_ratio": 0.20, "festival_lead_days": 7, "event_max_age_days": 7, "seasonal_max_age_days": 180, } def _p0_candidate(**overrides): candidate = { "title": "适合家庭分享的生活技巧", "publish_at": "2026-07-31T09:00:00+08:00", "duration_seconds": 30, "share_count": 1000, "content_50_plus_ratio": 0.20, "account_50_plus_ratio": 0.30, } candidate.update(overrides) return candidate def test_p0_gate_enforces_boundaries_and_time_context() -> None: assert evaluate_candidate_gate(_p0_candidate(), _P0_RULES)["primary_eligible"] short = evaluate_candidate_gate( _p0_candidate(duration_seconds=Decimal("29.999")), _P0_RULES, ) assert "DURATION_TOO_SHORT" in short["failed_reason_codes"] low_share = evaluate_candidate_gate(_p0_candidate(share_count=999), _P0_RULES) assert "SHARE_COUNT_TOO_LOW" in low_share["failed_reason_codes"] low_elder = evaluate_candidate_gate( _p0_candidate(content_50_plus_ratio=0.199), _P0_RULES, ) assert "CONTENT_50_PLUS_TOO_LOW" in low_elder["failed_reason_codes"] morning = evaluate_candidate_gate( _p0_candidate(title="早上好,送给家人的祝福"), _P0_RULES, ) assert "DAYPART_EXPIRED" in morning["failed_reason_codes"] festival = evaluate_candidate_gate( _p0_candidate(title="春节祝福送给全家"), _P0_RULES, ) assert "FESTIVAL_OUT_OF_WINDOW" in festival["failed_reason_codes"] def test_p0_gate_keeps_content_and_account_portraits_separate() -> None: conflict = evaluate_candidate_gate( _p0_candidate( content_50_plus_ratio=0.28, account_50_plus_ratio=0.08, ), _P0_RULES, ) assert conflict["content_portrait_status"] == "pass" assert conflict["account_portrait_status"] == "fail" assert conflict["portrait_conflict"] is True assert conflict["primary_eligible"] is True account_only = evaluate_candidate_gate( _p0_candidate( content_50_plus_ratio=None, account_50_plus_ratio=0.55, ), _P0_RULES, ) assert "CONTENT_PORTRAIT_MISSING" in account_only["failed_reason_codes"] def test_search_candidate_persists_publish_time_duration_and_shares() -> None: candidate = _candidate_from_search_result( { "aweme_id": "video-p0", "desc": "测试视频", "duration_ms": 65000, "publish_at": "2026-07-31T08:30:00+08:00", "statistics": {"share_count": 45}, }, "测试关键词", ) assert candidate is not None assert candidate["duration_seconds"] == Decimal("65.000") assert candidate["publish_at"] == datetime(2026, 7, 31, 8, 30) assert candidate["share_count"] == 45 def test_all_search_sources_default_to_thirty_seconds() -> None: assert INTERNAL_SEARCH_MIN_DURATION == 30 assert TIKHUB_SEARCH_MIN_DURATION == 30 assert AUTHOR_SEARCH_MIN_DURATION == 30 def test_repository_rejects_primary_that_fails_p0_gate() -> None: factory = _expire_on_commit_session_factory() with factory() as session: session.add( VideoDiscoveryRun( id=1, run_id="p0-gate-run", demand_word="生活技巧", relevant_points_json="[]", status="running", rule_version="test-p0", rule_config_json=json.dumps(_P0_RULES, ensure_ascii=False), ) ) session.add( VideoDiscoveryCandidate( id=2, run_id="p0-gate-run", aweme_id="video-p0", title="适合家庭分享的生活技巧", publish_at=datetime(2026, 7, 31, 9, 0), duration_seconds=Decimal("30.000"), share_count=999, content_50_plus_ratio=Decimal("0.280000"), account_50_plus_ratio=Decimal("0.080000"), decision_bucket="pending_evaluation", ) ) session.commit() with factory() as session: repo = VideoDiscoveryRepository(session) with pytest.raises(ValueError, match="SHARE_COUNT_TOO_LOW"): repo.update_candidates( "p0-gate-run", [{"candidate_id": 2, "decision_bucket": "primary"}], ) @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