test_discover_videos_from_demands.py 21 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692
  1. from __future__ import annotations
  2. import asyncio
  3. import importlib
  4. import json
  5. from contextlib import contextmanager
  6. from unittest.mock import patch
  7. import pytest
  8. from sqlalchemy import create_engine, event, select
  9. from sqlalchemy.orm import Session, sessionmaker
  10. from agents.find_agent import create_find_agent
  11. from agents.find_agent.async_runner import arun_find_agent
  12. from agents.find_agent.demand_run import (
  13. FindDemandContext,
  14. FindDemandPoint,
  15. FindDemandVideo,
  16. build_find_agent_user_input,
  17. prepare_video_discovery_run,
  18. )
  19. from agents.find_agent.tools import video_discovery_store
  20. from agents.find_agent.tools.batch_search_and_record import batch_search_and_record
  21. from supply_infra.db.models.video_discovery import (
  22. VideoDiscoveryCandidate,
  23. VideoDiscoveryRun,
  24. VideoDiscoverySearch,
  25. )
  26. from supply_infra.db.repositories.video_discovery_repo import (
  27. VideoDiscoveryRepository,
  28. )
  29. from supply_infra.scheduler.jobs.discover_videos_from_demands import (
  30. discover_videos_from_demands,
  31. )
  32. def _expire_on_commit_session_factory() -> sessionmaker[Session]:
  33. """Match production get_session: commit expires ORM instances."""
  34. engine = create_engine("sqlite+pysqlite:///:memory:")
  35. # SQLite 对 BigInteger PK 不会自增,测试里显式写入 id。
  36. VideoDiscoveryRun.__table__.create(engine)
  37. VideoDiscoveryCandidate.__table__.create(engine)
  38. return sessionmaker(bind=engine, autoflush=False, autocommit=False)
  39. def _patch_service_session(
  40. monkeypatch: pytest.MonkeyPatch,
  41. factory: sessionmaker[Session],
  42. ) -> None:
  43. @contextmanager
  44. def get_test_session():
  45. session = factory()
  46. try:
  47. yield session
  48. session.commit()
  49. except Exception:
  50. session.rollback()
  51. raise
  52. finally:
  53. session.close()
  54. monkeypatch.setattr(
  55. "supply_infra.services.video_discovery_service.get_session",
  56. get_test_session,
  57. )
  58. import supply_infra.services.video_discovery_service as service_module
  59. service_module._default_service = None
  60. def _seed_run(
  61. factory: sessionmaker[Session],
  62. *,
  63. run_id: str,
  64. demand_grade_id: int,
  65. status: str = "running",
  66. row_id: int = 1,
  67. ) -> None:
  68. with factory() as session:
  69. session.add(
  70. VideoDiscoveryRun(
  71. id=row_id,
  72. run_id=run_id,
  73. biz_dt="20260728",
  74. demand_grade_id=demand_grade_id,
  75. demand_word="广场舞",
  76. seed_video_id="vid-1",
  77. seed_video_title="参考标题",
  78. relevant_points_json="[]",
  79. status=status,
  80. )
  81. )
  82. session.commit()
  83. def test_create_video_discovery_run_reuses_precreated_run(
  84. monkeypatch: pytest.MonkeyPatch,
  85. ) -> None:
  86. """定时任务预创建 run 后,Agent 复用 run_id 不得触发 DetachedInstanceError。"""
  87. factory = _expire_on_commit_session_factory()
  88. _patch_service_session(monkeypatch, factory)
  89. _seed_run(factory, run_id="precreated-run", demand_grade_id=101)
  90. payload = json.loads(
  91. video_discovery_store.create_video_discovery_run(
  92. demand_word="广场舞",
  93. seed_video_title="参考标题",
  94. relevant_points=[],
  95. run_id="precreated-run",
  96. demand_grade_id=101,
  97. )
  98. )
  99. assert "error" not in payload
  100. assert payload["run_id"] == "precreated-run"
  101. assert payload["status"] == "running"
  102. assert payload["pre_created"] is True
  103. def test_prepare_then_reuse_run_id_scheduled_flow(
  104. monkeypatch: pytest.MonkeyPatch,
  105. ) -> None:
  106. """完整调度衔接:prepare 得到 run_id → create 复用,全程无 DetachedInstanceError。"""
  107. factory = _expire_on_commit_session_factory()
  108. _patch_service_session(monkeypatch, factory)
  109. # 先写入一条 failed 记录,prepare 会原地重置为 running 并复用 run_id。
  110. _seed_run(
  111. factory,
  112. run_id="scheduled-run",
  113. demand_grade_id=202,
  114. status="failed",
  115. )
  116. ctx = FindDemandContext(
  117. biz_dt="20260728",
  118. demand_grade_id=202,
  119. demand_name="广场舞",
  120. grade="S",
  121. videos=[
  122. FindDemandVideo(
  123. video_id="vid-1",
  124. title="参考标题",
  125. points=[
  126. FindDemandPoint(
  127. point="动作简单",
  128. point_type="key",
  129. )
  130. ],
  131. )
  132. ],
  133. )
  134. run_id, skip_reason = prepare_video_discovery_run(ctx)
  135. assert skip_reason is None
  136. assert run_id == "scheduled-run"
  137. reuse_id, reuse_skip = prepare_video_discovery_run(ctx)
  138. assert reuse_id == "scheduled-run"
  139. assert reuse_skip is None
  140. payload = json.loads(
  141. video_discovery_store.create_video_discovery_run(
  142. demand_word=ctx.demand_name,
  143. seed_video_title="参考标题",
  144. relevant_points=[{"point": "动作简单", "point_type": "key"}],
  145. run_id=run_id,
  146. demand_grade_id=ctx.demand_grade_id,
  147. seed_video_id="vid-1",
  148. )
  149. )
  150. assert "error" not in payload
  151. assert payload["pre_created"] is True
  152. assert payload["run_id"] == run_id
  153. def test_prepare_skips_when_run_already_finished(
  154. monkeypatch: pytest.MonkeyPatch,
  155. ) -> None:
  156. factory = _expire_on_commit_session_factory()
  157. _patch_service_session(monkeypatch, factory)
  158. _seed_run(
  159. factory,
  160. run_id="finished-run",
  161. demand_grade_id=303,
  162. status="finished",
  163. )
  164. ctx = FindDemandContext(
  165. biz_dt="20260728",
  166. demand_grade_id=303,
  167. demand_name="广场舞",
  168. grade="S",
  169. videos=[
  170. FindDemandVideo(
  171. video_id="vid-1",
  172. title="参考标题",
  173. points=[FindDemandPoint(point="动作简单", point_type="key")],
  174. )
  175. ],
  176. )
  177. run_id, skip_reason = prepare_video_discovery_run(ctx)
  178. assert run_id is None
  179. assert skip_reason is not None
  180. assert "finished" in skip_reason
  181. def test_video_discovery_models_exclude_unused_columns() -> None:
  182. run_columns = set(VideoDiscoveryRun.__table__.columns.keys())
  183. candidate_columns = set(VideoDiscoveryCandidate.__table__.columns.keys())
  184. assert "backup_count" not in run_columns
  185. assert {
  186. "video_url",
  187. "content_analysis",
  188. "content_analysis_verified",
  189. "hit_points_json",
  190. "publish_timestamp",
  191. "detail_verified",
  192. "content_portrait_attempted",
  193. "account_portrait_attempted",
  194. "age_portraits_normalized",
  195. "expansion_worthy_tags_json",
  196. "confidence",
  197. "relevance_reason",
  198. "elder_reason",
  199. "share_reason",
  200. "manual_review_note",
  201. "manual_review_status",
  202. }.isdisjoint(candidate_columns)
  203. assert "search_id" in candidate_columns
  204. candidate_constraints = {
  205. constraint.name
  206. for constraint in VideoDiscoveryCandidate.__table__.constraints
  207. }
  208. search_constraints = {
  209. constraint.name
  210. for constraint in VideoDiscoverySearch.__table__.constraints
  211. }
  212. assert "uk_video_discovery_candidate_run_aweme" not in candidate_constraints
  213. assert "fk_video_discovery_candidate_search" in candidate_constraints
  214. assert "uk_video_discovery_search_key" not in search_constraints
  215. def test_create_find_agent_registers_discovery_tools() -> None:
  216. agent = create_find_agent()
  217. assert agent.name == "find_agent"
  218. assert "batch_search_and_record" in agent.tools.list_tools()
  219. assert "batch_update_video_discovery_candidates" in agent.tools.list_tools()
  220. assert "update_video_discovery_run_status" in agent.tools.list_tools()
  221. assert "create_video_discovery_run" not in agent.tools.list_tools()
  222. assert "batch_save_video_candidate_evaluations" not in agent.tools.list_tools()
  223. assert "audit_video_discovery_run" not in agent.tools.list_tools()
  224. assert "query_video_discovery_state" in agent.tools.list_tools()
  225. def test_find_agent_input_uses_reference_videos_without_seed_fields() -> None:
  226. ctx = FindDemandContext(
  227. biz_dt="20260729",
  228. demand_grade_id=101,
  229. demand_name="广场舞",
  230. grade="S",
  231. videos=[
  232. FindDemandVideo(
  233. video_id="vid-1",
  234. title="参考标题",
  235. points=[FindDemandPoint(point="动作简单", point_type="key")],
  236. )
  237. ],
  238. )
  239. user_input = build_find_agent_user_input(ctx, "scheduled-run")
  240. assert "seed_video_id:" not in user_input
  241. assert "seed_video_title:" not in user_input
  242. assert "reference_videos:" in user_input
  243. assert '"video_id": "vid-1"' in user_input
  244. assert '"title": "参考标题"' in user_input
  245. assert "create_video_discovery_run" not in user_input
  246. assert "relevant_points" not in user_input
  247. @pytest.mark.asyncio
  248. async def test_douyin_search_automatically_persists_page(
  249. monkeypatch: pytest.MonkeyPatch,
  250. ) -> None:
  251. search_module = importlib.import_module(
  252. "agents.find_agent.tools.douyin_search"
  253. )
  254. async def fake_raw_search(**_kwargs):
  255. return json.dumps(
  256. {
  257. "results_count": 1,
  258. "has_more": False,
  259. "search_results": [{"aweme_id": "auto-saved"}],
  260. }
  261. )
  262. persisted: dict[str, object] = {}
  263. def fake_persist(payload_json: str, **kwargs):
  264. persisted.update(kwargs)
  265. payload = json.loads(payload_json)
  266. payload.update(
  267. {
  268. "persisted": True,
  269. "search_id": 11,
  270. "new_candidate_count": 1,
  271. "candidates": [
  272. {
  273. "candidate_id": 21,
  274. "search_id": 11,
  275. "aweme_id": "auto-saved",
  276. "title": "自动保存",
  277. "decision_bucket": "pending_evaluation",
  278. }
  279. ],
  280. }
  281. )
  282. return json.dumps(payload)
  283. monkeypatch.setattr(search_module, "_douyin_search_raw", fake_raw_search)
  284. monkeypatch.setattr(search_module, "persist_search_payload", fake_persist)
  285. result = json.loads(
  286. await search_module.douyin_search(
  287. run_id="run-auto-save",
  288. keyword="广场舞",
  289. query_reason="验证需求根搜索",
  290. source_type="demand",
  291. )
  292. )
  293. assert result["persisted"] is True
  294. assert result["search_id"] == 11
  295. assert result["candidates"][0]["candidate_id"] == 21
  296. assert persisted["run_id"] == "run-auto-save"
  297. assert persisted["keyword"] == "广场舞"
  298. assert persisted["provider"] == "internal_keyword"
  299. @pytest.mark.asyncio
  300. async def test_batch_search_records_each_page_and_carries_parent(
  301. monkeypatch: pytest.MonkeyPatch,
  302. ) -> None:
  303. batch_module = importlib.import_module(
  304. "agents.find_agent.tools.batch_search_and_record"
  305. )
  306. calls: list[dict[str, object]] = []
  307. async def fake_search(**kwargs):
  308. calls.append(kwargs)
  309. page_no = int(kwargs["page_no"])
  310. return json.dumps(
  311. {
  312. "results_count": 1,
  313. "has_more": page_no == 1,
  314. "next_cursor": "next-page" if page_no == 1 else None,
  315. "persisted": True,
  316. "search_id": 100 + page_no,
  317. "new_candidate_count": 1,
  318. "candidates": [
  319. {
  320. "candidate_id": 200 + page_no,
  321. "search_id": 100 + page_no,
  322. "aweme_id": "same-video",
  323. "title": f"第 {page_no} 页",
  324. "decision_bucket": "pending_evaluation",
  325. }
  326. ],
  327. }
  328. )
  329. monkeypatch.setattr(batch_module, "douyin_search", fake_search)
  330. result = json.loads(
  331. await batch_search_and_record(
  332. run_id="run-batch",
  333. searches=[
  334. {
  335. "keyword": "广场舞",
  336. "query_reason": "验证需求根搜索",
  337. "source_type": "demand",
  338. "max_pages": 2,
  339. }
  340. ],
  341. )
  342. )
  343. assert result["saved_page_count"] == 2
  344. assert result["new_candidate_count"] == 2
  345. assert result["tasks"][0]["pages"][0]["candidates"][0]["candidate_id"] == 201
  346. assert result["tasks"][0]["pages"][1]["candidates"][0]["candidate_id"] == 202
  347. assert calls[0]["parent_search_id"] is None
  348. assert calls[1]["parent_search_id"] == 101
  349. assert calls[1]["cursor"] == "next-page"
  350. def test_batch_update_candidates_uses_database_candidate_id(
  351. monkeypatch: pytest.MonkeyPatch,
  352. ) -> None:
  353. captured: dict[str, object] = {}
  354. class FakeService:
  355. def update_candidates(self, run_id, rows):
  356. captured["run_id"] = run_id
  357. captured["rows"] = rows
  358. return {
  359. "updated_count": 1,
  360. "candidates": [
  361. {
  362. "candidate_id": 901,
  363. "search_id": 801,
  364. "aweme_id": "same-video",
  365. "decision_bucket": "primary",
  366. }
  367. ],
  368. }
  369. monkeypatch.setattr(
  370. video_discovery_store,
  371. "get_video_discovery_service",
  372. lambda: FakeService(),
  373. )
  374. result = json.loads(
  375. video_discovery_store.batch_update_video_discovery_candidates(
  376. run_id="run-update",
  377. items=[
  378. {
  379. "candidate_id": 901,
  380. "decision_bucket": "primary",
  381. "relevance_score": 0.8,
  382. "elder_score": 0.7,
  383. "share_score": 0.6,
  384. }
  385. ],
  386. )
  387. )
  388. assert result["updated_count"] == 1
  389. assert captured["run_id"] == "run-update"
  390. assert captured["rows"][0]["candidate_id"] == 901
  391. assert "aweme_id" not in captured["rows"][0]
  392. def test_each_search_inserts_new_candidate_occurrences() -> None:
  393. engine = create_engine("sqlite+pysqlite:///:memory:")
  394. VideoDiscoveryRun.__table__.create(engine)
  395. VideoDiscoverySearch.__table__.create(engine)
  396. VideoDiscoveryCandidate.__table__.create(engine)
  397. factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
  398. ids = {"search": 100, "candidate": 1000}
  399. @event.listens_for(factory.class_, "before_flush")
  400. def assign_sqlite_bigint_ids(session, _flush_context, _instances):
  401. for entity in session.new:
  402. if isinstance(entity, VideoDiscoverySearch) and entity.id is None:
  403. ids["search"] += 1
  404. entity.id = ids["search"]
  405. elif isinstance(entity, VideoDiscoveryCandidate) and entity.id is None:
  406. ids["candidate"] += 1
  407. entity.id = ids["candidate"]
  408. with factory() as session:
  409. session.add(
  410. VideoDiscoveryRun(
  411. id=1,
  412. run_id="run-occurrences",
  413. demand_word="广场舞",
  414. relevant_points_json="[]",
  415. status="running",
  416. )
  417. )
  418. session.commit()
  419. search_values = {
  420. "run_id": "run-occurrences",
  421. "search_key": "same-search-key",
  422. "keyword": "广场舞",
  423. "query_reason": "验证相同搜索也生成新记录",
  424. "source_type": "demand",
  425. "provider": "internal_keyword",
  426. "content_type": "视频",
  427. "sort_type": "综合排序",
  428. "publish_time": "不限",
  429. "cursor": "0",
  430. "page_no": 1,
  431. "results_count": 1,
  432. "new_candidate_count": 0,
  433. "has_more": 0,
  434. "status": "success",
  435. }
  436. candidate_rows = [
  437. {
  438. "aweme_id": "same-video",
  439. "title": "同一个视频",
  440. "_source_keyword": "广场舞",
  441. }
  442. ]
  443. with factory() as session:
  444. repo = VideoDiscoveryRepository(session)
  445. first_search, first_candidates = repo.save_search_page(
  446. dict(search_values),
  447. candidate_rows,
  448. )
  449. second_search, second_candidates = repo.save_search_page(
  450. dict(search_values),
  451. candidate_rows,
  452. )
  453. session.commit()
  454. assert first_search.id != second_search.id
  455. assert first_candidates[0].id != second_candidates[0].id
  456. assert first_candidates[0].search_id == first_search.id
  457. assert second_candidates[0].search_id == second_search.id
  458. with factory() as session:
  459. searches = session.scalars(select(VideoDiscoverySearch)).all()
  460. candidates = session.scalars(select(VideoDiscoveryCandidate)).all()
  461. assert len(searches) == 2
  462. assert len(candidates) == 2
  463. assert {candidate.aweme_id for candidate in candidates} == {"same-video"}
  464. def _seed_candidate(
  465. factory: sessionmaker[Session],
  466. *,
  467. run_id: str,
  468. aweme_id: str = "7631830155522179258",
  469. row_id: int = 1,
  470. ) -> None:
  471. with factory() as session:
  472. session.add(
  473. VideoDiscoveryCandidate(
  474. id=row_id,
  475. run_id=run_id,
  476. aweme_id=aweme_id,
  477. decision_bucket="pending_evaluation",
  478. )
  479. )
  480. session.commit()
  481. def test_list_skip_grade_ids_when_candidates_exist(
  482. monkeypatch: pytest.MonkeyPatch,
  483. ) -> None:
  484. factory = _expire_on_commit_session_factory()
  485. _patch_service_session(monkeypatch, factory)
  486. _seed_run(factory, run_id="running-empty", demand_grade_id=301, status="running")
  487. _seed_run(
  488. factory,
  489. run_id="running-with-candidates",
  490. demand_grade_id=302,
  491. status="running",
  492. row_id=2,
  493. )
  494. _seed_candidate(factory, run_id="running-with-candidates")
  495. from supply_infra.services.video_discovery_service import get_video_discovery_service
  496. skip_ids = get_video_discovery_service().list_skip_grade_ids("20260728")
  497. assert skip_ids == {302}
  498. def test_evaluate_find_agent_run_succeeds_when_candidates_exist(
  499. monkeypatch: pytest.MonkeyPatch,
  500. ) -> None:
  501. from agents.find_agent.run_outcome import evaluate_find_agent_run
  502. from supply_agent.types import AgentResult
  503. monkeypatch.setattr(
  504. "agents.find_agent.run_outcome.get_video_discovery_service",
  505. lambda: type(
  506. "Svc",
  507. (),
  508. {"has_candidates": staticmethod(lambda _run_id: True)},
  509. )(),
  510. )
  511. outcome = evaluate_find_agent_run(
  512. "run-1",
  513. AgentResult(content="任意文案", messages=[], iterations=3, tool_calls_made=0),
  514. )
  515. assert outcome.succeeded is True
  516. assert outcome.failure_reason is None
  517. def test_evaluate_find_agent_run_fails_without_candidates(
  518. monkeypatch: pytest.MonkeyPatch,
  519. ) -> None:
  520. from agents.find_agent.run_outcome import evaluate_find_agent_run
  521. from supply_agent.types import AgentResult
  522. monkeypatch.setattr(
  523. "agents.find_agent.run_outcome.get_video_discovery_service",
  524. lambda: type(
  525. "Svc",
  526. (),
  527. {"has_candidates": staticmethod(lambda _run_id: False)},
  528. )(),
  529. )
  530. outcome = evaluate_find_agent_run(
  531. "run-1",
  532. AgentResult(
  533. content="任务未完成(工具故障)",
  534. messages=[],
  535. iterations=1,
  536. tool_calls_made=0,
  537. ),
  538. )
  539. assert outcome.succeeded is False
  540. assert outcome.failure_reason == "no_candidates"
  541. @patch(
  542. "supply_infra.scheduler.jobs.discover_videos_from_demands.process_single_discover"
  543. )
  544. @patch("supply_infra.scheduler.jobs.discover_videos_from_demands._count_passed_videos")
  545. @patch(
  546. "supply_infra.scheduler.jobs.discover_videos_from_demands.filter_pending_contexts"
  547. )
  548. @patch(
  549. "supply_infra.scheduler.jobs.discover_videos_from_demands.list_find_demand_contexts"
  550. )
  551. def test_stops_discovery_after_200_passed_videos(
  552. mock_list_contexts,
  553. mock_filter_contexts,
  554. mock_count_passed,
  555. mock_process,
  556. ) -> None:
  557. contexts = [
  558. FindDemandContext(
  559. biz_dt="20260727",
  560. demand_grade_id=1,
  561. demand_name="需求A",
  562. grade="S",
  563. ),
  564. FindDemandContext(
  565. biz_dt="20260727",
  566. demand_grade_id=2,
  567. demand_name="需求B",
  568. grade="A",
  569. ),
  570. ]
  571. mock_list_contexts.return_value = ("20260727", contexts)
  572. mock_filter_contexts.return_value = (
  573. contexts,
  574. {"total_loaded": 2, "skipped_already_done": 0},
  575. )
  576. mock_count_passed.side_effect = [199, 200]
  577. mock_process.return_value = {"success": True, "skipped": False}
  578. result = discover_videos_from_demands("20260727", workers=1)
  579. assert mock_process.call_count == 1
  580. assert result["processed"] == 1
  581. assert result["passed_videos"] == 200
  582. assert result["stopped_by_passed_video_limit"] is True
  583. class _AsyncClient:
  584. def __init__(self) -> None:
  585. self.closed = False
  586. async def close(self) -> None:
  587. self.closed = True
  588. class _SlowAgent:
  589. def __init__(self) -> None:
  590. self.llm = type("LLM", (), {"_async_client": _AsyncClient()})()
  591. async def arun_core(self, _user_input: str) -> None:
  592. await asyncio.sleep(60)
  593. @pytest.mark.asyncio
  594. async def test_find_agent_timeout_closes_async_client() -> None:
  595. agent = _SlowAgent()
  596. with pytest.raises(TimeoutError, match="find_agent timed out"):
  597. await arun_find_agent(
  598. agent, # type: ignore[arg-type]
  599. "test",
  600. run_id="test-run",
  601. timeout_seconds=0.01,
  602. )
  603. assert agent.llm._async_client.closed is True