test_find_agent_v2.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477
  1. from __future__ import annotations
  2. import json
  3. from dataclasses import replace
  4. from pathlib import Path
  5. from types import SimpleNamespace
  6. import pytest
  7. from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
  8. from langchain_core.messages import AIMessage
  9. from find_agent_v2.graph import FindAgentRoundGraph
  10. from find_agent_v2.runtime import DelegateArgs, FindAgentNodeHost
  11. from find_agent_v2.demand_context import (
  12. V2DemandContext,
  13. V2ReferencePoint,
  14. V2ReferenceVideo,
  15. _points_from_expansions,
  16. build_v2_user_input,
  17. )
  18. from find_agent_v2.gates import build_rule_snapshot, evaluate_candidate_gate
  19. from find_agent_v2.models import (
  20. FindAgentV2Candidate,
  21. FindAgentV2Evidence,
  22. FindAgentV2Round,
  23. FindAgentV2Run,
  24. FindAgentV2Search,
  25. )
  26. from find_agent_v2.observability import (
  27. GRAPH_SPEC,
  28. MODULE_TITLES,
  29. OBAGENT_AGENT,
  30. OBAGENT_PROJECT,
  31. OBAGENT_ROUND_ANCHOR,
  32. NullObserver,
  33. )
  34. from find_agent_v2.prompts import COMMON_RULES
  35. from find_agent_v2.providers import normalize_age_pair
  36. from find_agent_v2.state import DiscoverySnapshot, FindAgentState, NodeRun
  37. from find_agent_v2.tools import (
  38. EVALUATION_TOOLS,
  39. EVIDENCE_TOOLS,
  40. REPORT_TOOLS,
  41. SEARCH_TOOLS,
  42. )
  43. def _names(functions) -> set[str]:
  44. return {getattr(fn, "_tool_name", fn.__name__) for fn in functions}
  45. def test_v2_orm_uses_only_new_table_namespace() -> None:
  46. assert {
  47. FindAgentV2Run.__tablename__,
  48. FindAgentV2Round.__tablename__,
  49. FindAgentV2Search.__tablename__,
  50. FindAgentV2Candidate.__tablename__,
  51. FindAgentV2Evidence.__tablename__,
  52. } == {
  53. "find_agent_v2_run",
  54. "find_agent_v2_round",
  55. "find_agent_v2_search",
  56. "find_agent_v2_candidate",
  57. "find_agent_v2_evidence",
  58. }
  59. assert "obagent_run_uid" in FindAgentV2Run.__table__.columns
  60. assert {"input_tokens", "output_tokens", "total_tokens", "cost_usd"} <= {
  61. column.name for column in FindAgentV2Run.__table__.columns
  62. }
  63. assert not any(table.foreign_key_constraints for table in (
  64. FindAgentV2Run.__table__,
  65. FindAgentV2Round.__table__,
  66. FindAgentV2Search.__table__,
  67. FindAgentV2Candidate.__table__,
  68. FindAgentV2Evidence.__table__,
  69. ))
  70. def test_v2_package_has_no_legacy_business_imports() -> None:
  71. package = Path(__file__).parents[2] / "find_agent_v2"
  72. source = "\n".join(path.read_text(encoding="utf-8") for path in package.glob("*.py"))
  73. assert "agents.find_agent" not in source
  74. assert "video_discovery_gates" not in source
  75. assert "services.video_discovery" not in source
  76. assert "system_prompt.md" not in source
  77. def test_project_orm_defines_no_database_foreign_keys() -> None:
  78. import supply_infra.db.models # noqa: F401
  79. from supply_infra.db.base import Base
  80. assert not {
  81. table.name: sorted(fk.name or "<unnamed>" for fk in table.foreign_key_constraints)
  82. for table in Base.metadata.tables.values()
  83. if table.foreign_key_constraints
  84. }
  85. def test_v2_owns_age_portrait_normalization() -> None:
  86. normalized = normalize_age_pair(
  87. {"年龄": {"50-": {"percentage": "35%", "preference": 120}}},
  88. {"年龄": {"50岁以上": {"percentage": 0.25, "preference": 110}}},
  89. )
  90. assert normalized["content"]["older_ratio"] == 0.35
  91. assert normalized["account"]["older_ratio"] == 0.25
  92. assert normalized["consistency"] == "aligned"
  93. def test_v2_owns_primary_candidate_gate() -> None:
  94. rules = build_rule_snapshot()
  95. candidate = {
  96. "title": "适合父母的实用生活建议",
  97. "publish_at": rules["current_datetime"],
  98. "duration_seconds": 60,
  99. "share_count": 2000,
  100. "content_50_plus_ratio": 0.35,
  101. "account_50_plus_ratio": 0.25,
  102. "relevance_score": 0.8,
  103. "elder_score": 0.8,
  104. "share_score": 0.8,
  105. "value_score": 0.8,
  106. }
  107. result = evaluate_candidate_gate(candidate, rules)
  108. assert result["status"] == "pass"
  109. assert result["primary_eligible"] is True
  110. def test_v2_gate_rejects_explicitly_low_evidence() -> None:
  111. rules = build_rule_snapshot()
  112. candidate = {
  113. "title": "普通视频",
  114. "publish_at": rules["current_datetime"],
  115. "duration_seconds": 10,
  116. "share_count": 3,
  117. "content_50_plus_ratio": 0.01,
  118. "account_50_plus_ratio": 0.02,
  119. }
  120. result = evaluate_candidate_gate(candidate, rules)
  121. assert result["status"] == "fail"
  122. assert {"DURATION_TOO_SHORT", "SHARE_COUNT_TOO_LOW", "PORTRAIT_50_PLUS_TOO_LOW"} <= set(
  123. result["failed_reason_codes"]
  124. )
  125. def test_v2_context_deduplicates_expansion_points() -> None:
  126. rows = [
  127. SimpleNamespace(
  128. video_id="v1", point_type="purpose", expanded_text="照顾父母",
  129. point_desc="描述一",
  130. ),
  131. SimpleNamespace(
  132. video_id="v1", point_type="purpose", expanded_text="照顾父母",
  133. point_desc="重复描述",
  134. ),
  135. SimpleNamespace(
  136. video_id="v1", point_type="invalid", expanded_text="忽略",
  137. point_desc=None,
  138. ),
  139. ]
  140. order, points = _points_from_expansions(rows)
  141. assert order == ["v1"]
  142. assert len(points["v1"]) == 1
  143. assert points["v1"][0].point == "照顾父母"
  144. def test_v2_context_builds_self_contained_user_input() -> None:
  145. context = V2DemandContext(
  146. biz_dt="20260811",
  147. demand_grade_id=123,
  148. demand_name="测试需求",
  149. grade="S",
  150. score=95.0,
  151. videos=[V2ReferenceVideo(
  152. video_id="video-1",
  153. title="参考视频",
  154. points=[V2ReferencePoint("关键内容", "key", "关键描述")],
  155. )],
  156. )
  157. raw = build_v2_user_input(
  158. context,
  159. run_id="v2-test-run",
  160. rules={"rule_version": "v2-test"},
  161. )
  162. assert '"run_id": "v2-test-run"' in raw
  163. assert '"demand_grade_id": 123' in raw
  164. assert '"reference_videos"' in raw
  165. assert '"关键内容"' in raw
  166. def test_obagent_identity_and_round_structure_are_stable() -> None:
  167. assert OBAGENT_PROJECT == "find_agent_v2"
  168. assert OBAGENT_AGENT == "find_agent_v2"
  169. assert OBAGENT_ROUND_ANCHOR == {"in": "run", "on": ["graph"]}
  170. assert [node["key"] for node in GRAPH_SPEC["nodes"]] == [
  171. "supervisor", "search", "evidence", "evaluator",
  172. ]
  173. assert set(MODULE_TITLES) == {"supervisor", "search", "evidence", "evaluator", "report"}
  174. def test_round_graph_is_a_compiled_langgraph() -> None:
  175. service = _FakeService(pending_after_search=0)
  176. graph = FindAgentRoundGraph(service=service, runner=_FakeRunner(service))
  177. drawable = graph.app.get_graph()
  178. assert {"supervisor", "search", "evidence", "evaluator"} <= set(drawable.nodes)
  179. assert graph.obagent_spec.get("nodes")
  180. def test_runtime_exposes_bounded_delegate_schema() -> None:
  181. field = DelegateArgs.model_fields["requests"]
  182. assert field.metadata
  183. assert FindAgentNodeHost.__module__ == "find_agent_v2.runtime"
  184. def test_runtime_usage_can_be_reset_between_resume_attempts() -> None:
  185. host = FindAgentNodeHost(observer=NullObserver())
  186. host.usage["total_tokens"] = 99
  187. host.reset_usage()
  188. assert host.usage == {
  189. "input_tokens": 0,
  190. "output_tokens": 0,
  191. "total_tokens": 0,
  192. "cost": 0.0,
  193. }
  194. @pytest.mark.asyncio
  195. async def test_langchain_runtime_runs_without_network(monkeypatch) -> None:
  196. host = FindAgentNodeHost(observer=NullObserver())
  197. fake_model = FakeMessagesListChatModel(responses=[AIMessage(content="done")])
  198. monkeypatch.setattr(host, "_model", lambda _role: fake_model)
  199. result = await host.run_node(
  200. node="supervisor",
  201. round_index=1,
  202. system_prompt="plan",
  203. user_content="task",
  204. tools=(),
  205. max_iterations=2,
  206. allow_delegation=False,
  207. )
  208. assert result.content == "done"
  209. assert result.iterations == 1
  210. assert result.tool_calls_made == 0
  211. def test_stage_tool_allowlists_are_physical_and_isolated() -> None:
  212. assert _names(SEARCH_TOOLS) == {"search_videos_v2", "query_find_agent_v2_state"}
  213. assert _names(EVIDENCE_TOOLS) == {
  214. "fetch_candidate_details_v2",
  215. "fetch_candidate_portraits_v2",
  216. "query_pending_candidates_v2",
  217. }
  218. assert _names(EVALUATION_TOOLS) == {
  219. "evaluate_candidates_v2",
  220. "query_pending_candidates_v2",
  221. }
  222. assert _names(REPORT_TOOLS) == {"query_find_agent_v2_state"}
  223. all_names = _names((*SEARCH_TOOLS, *EVIDENCE_TOOLS, *EVALUATION_TOOLS, *REPORT_TOOLS))
  224. assert not any(name.startswith("batch_update_video_discovery") for name in all_names)
  225. assert "query_video_discovery_state" not in all_names
  226. def test_common_prompt_points_to_v2_tables_and_tools() -> None:
  227. assert "find_agent_v2_run" in COMMON_RULES
  228. assert "video_discovery_run" not in COMMON_RULES
  229. assert "batch_search_and_record" not in COMMON_RULES
  230. assert "batch_update_video_discovery_candidates" not in COMMON_RULES
  231. class _FakeService:
  232. def __init__(self, *, pending_after_search: int) -> None:
  233. self.pending_after_search = pending_after_search
  234. self.stage = "start"
  235. self.updates: list[dict] = []
  236. def get_full_state(self, run_id: str, **_kwargs):
  237. snapshot = self.snapshot(run_id)
  238. evidence_status = "pending" if self.stage in {"start", "searched"} else "success"
  239. return {
  240. "run": {"run_id": run_id},
  241. "searches": [],
  242. "candidates": [
  243. {
  244. "candidate_id": index,
  245. "decision_bucket": "pending_evaluation",
  246. "detail_status": evidence_status,
  247. "portrait_status": evidence_status,
  248. }
  249. for index in range(1, snapshot.pending_count + 1)
  250. ],
  251. }
  252. def snapshot(self, _run_id: str) -> DiscoverySnapshot:
  253. base = DiscoverySnapshot("running", 0, 0, 0, 0, 0, 0)
  254. if self.stage == "searched":
  255. return replace(base, search_count=1, candidate_count=self.pending_after_search,
  256. pending_count=self.pending_after_search)
  257. if self.stage in {"evidenced", "batched"}:
  258. return replace(base, search_count=1, candidate_count=self.pending_after_search,
  259. pending_count=self.pending_after_search)
  260. if self.stage == "evaluated":
  261. return replace(base, search_count=1, candidate_count=self.pending_after_search,
  262. rejected_count=self.pending_after_search)
  263. return base
  264. def update_round(self, _run_id: str, _round_index: int, **kwargs) -> None:
  265. self.updates.append(kwargs)
  266. def recount_valid_primary(self, _run_id: str) -> int:
  267. return 0
  268. class _FakeRunner:
  269. def __init__(self, service: _FakeService) -> None:
  270. self.service = service
  271. self.calls: list[tuple[str, set[str]]] = []
  272. async def run_node(self, *, node, round_index, tools=(), **_kwargs) -> NodeRun:
  273. self.calls.append((node, _names(tools)))
  274. if node == "search":
  275. self.service.stage = "searched"
  276. elif node == "evidence":
  277. self.service.stage = "evidenced"
  278. elif node == "evaluator":
  279. self.service.stage = "evaluated"
  280. if node == "supervisor":
  281. if self.service.stage == "evaluated":
  282. next_action = "finish"
  283. elif self.service.stage in {"evidenced", "batched"}:
  284. next_action = "evaluator"
  285. elif self.service.stage == "searched" and not self.service.pending_after_search:
  286. next_action = "finish"
  287. else:
  288. next_action = "search"
  289. content = json.dumps({
  290. "next_action": next_action,
  291. "worker_count": 4,
  292. "evidence_scope": "both",
  293. "plan": {"searches": []},
  294. })
  295. else:
  296. content = '{"searches": []}'
  297. return NodeRun(node, round_index, content, 1, 0)
  298. @pytest.mark.asyncio
  299. async def test_round_graph_supervisor_routes_with_guarded_allowlists() -> None:
  300. service = _FakeService(pending_after_search=2)
  301. runner = _FakeRunner(service)
  302. graph = FindAgentRoundGraph(service=service, runner=runner)
  303. state = FindAgentState(run_id="new-run", user_input="task", round_index=1)
  304. result = await graph.invoke(state)
  305. assert [name for name, _ in runner.calls] == [
  306. "supervisor", "search", "supervisor", "evidence", "evidence",
  307. "supervisor", "evaluator", "supervisor",
  308. ]
  309. assert runner.calls[0][1] == set()
  310. assert runner.calls[1][1] == _names(SEARCH_TOOLS)
  311. assert runner.calls[2][1] == set()
  312. assert runner.calls[3][1] <= _names(EVIDENCE_TOOLS)
  313. assert runner.calls[4][1] <= _names(EVIDENCE_TOOLS)
  314. assert runner.calls[6][1] == _names(EVALUATION_TOOLS)
  315. assert result.phase == "done"
  316. assert result.snapshot is not None and result.snapshot.pending_count == 0
  317. @pytest.mark.asyncio
  318. async def test_round_graph_skips_evidence_and_evaluation_without_candidates() -> None:
  319. service = _FakeService(pending_after_search=0)
  320. runner = _FakeRunner(service)
  321. graph = FindAgentRoundGraph(service=service, runner=runner)
  322. await graph.invoke(FindAgentState(run_id="new-run", user_input="task", round_index=1))
  323. assert [name for name, _ in runner.calls] == [
  324. "supervisor", "search", "supervisor",
  325. ]
  326. @pytest.mark.asyncio
  327. async def test_round_graph_reenters_evaluator_until_pending_queue_is_empty() -> None:
  328. service = _FakeService(pending_after_search=3)
  329. class BatchedRunner(_FakeRunner):
  330. def __init__(self, fake_service: _FakeService) -> None:
  331. super().__init__(fake_service)
  332. self.remaining = 3
  333. async def run_node(self, *, node, round_index, tools=(), **kwargs) -> NodeRun:
  334. result = await super().run_node(
  335. node=node, round_index=round_index, tools=tools, **kwargs,
  336. )
  337. if node == "evaluator":
  338. self.remaining -= 1
  339. self.service.stage = "evaluated" if self.remaining == 0 else "batched"
  340. return result
  341. runner = BatchedRunner(service)
  342. original_snapshot = service.snapshot
  343. def batched_snapshot(run_id: str) -> DiscoverySnapshot:
  344. if service.stage == "batched":
  345. base = original_snapshot(run_id)
  346. return replace(
  347. base,
  348. search_count=1,
  349. candidate_count=3,
  350. pending_count=runner.remaining,
  351. rejected_count=3 - runner.remaining,
  352. )
  353. return original_snapshot(run_id)
  354. service.snapshot = batched_snapshot # type: ignore[method-assign]
  355. graph = FindAgentRoundGraph(service=service, runner=runner)
  356. result = await graph.invoke(
  357. FindAgentState(run_id="batched-run", user_input="task", round_index=1),
  358. )
  359. assert [name for name, _ in runner.calls].count("evaluator") == 3
  360. assert result.snapshot is not None and result.snapshot.pending_count == 0
  361. def test_supervisor_policy_overrides_unsafe_finish_and_clamps_workers() -> None:
  362. service = _FakeService(pending_after_search=2)
  363. service.stage = "searched"
  364. graph = FindAgentRoundGraph(service=service, runner=_FakeRunner(service))
  365. action, _reason, workers, scope = graph._approve_action(
  366. {"run_id": "r", "search_actions": 1, "action_count": 1},
  367. {"next_action": "finish", "worker_count": 99, "evidence_scope": "invalid"},
  368. )
  369. assert action == "evidence"
  370. assert workers == 8
  371. assert scope == "both"
  372. def test_supervisor_can_choose_an_extra_search_within_budget() -> None:
  373. service = _FakeService(pending_after_search=0)
  374. service.stage = "searched"
  375. graph = FindAgentRoundGraph(service=service, runner=_FakeRunner(service))
  376. action, *_ = graph._approve_action(
  377. {"run_id": "r", "search_actions": 1, "action_count": 1},
  378. {"next_action": "search", "worker_count": 2},
  379. )
  380. assert action == "search"
  381. @pytest.mark.asyncio
  382. async def test_round_graph_retries_one_stagnant_evaluator_response() -> None:
  383. service = _FakeService(pending_after_search=2)
  384. class RetryRunner(_FakeRunner):
  385. evaluator_calls = 0
  386. async def run_node(self, *, node, round_index, tools=(), **kwargs) -> NodeRun:
  387. if node == "evaluator":
  388. self.calls.append((node, _names(tools)))
  389. self.evaluator_calls += 1
  390. if self.evaluator_calls == 2:
  391. self.service.stage = "evaluated"
  392. return NodeRun(node, round_index, "", 1, 0)
  393. return await super().run_node(
  394. node=node, round_index=round_index, tools=tools, **kwargs,
  395. )
  396. runner = RetryRunner(service)
  397. graph = FindAgentRoundGraph(service=service, runner=runner)
  398. result = await graph.invoke(
  399. FindAgentState(run_id="retry-run", user_input="task", round_index=1),
  400. )
  401. assert runner.evaluator_calls == 2
  402. assert result.snapshot is not None and result.snapshot.pending_count == 0