test_find_agent_v2.py 17 KB

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