| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477 |
- from __future__ import annotations
- import json
- from dataclasses import replace
- from pathlib import Path
- from types import SimpleNamespace
- import pytest
- from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
- from langchain_core.messages import AIMessage
- from find_agent_v2.graph import FindAgentRoundGraph
- from find_agent_v2.runtime import DelegateArgs, FindAgentNodeHost
- from find_agent_v2.demand_context import (
- V2DemandContext,
- V2ReferencePoint,
- V2ReferenceVideo,
- _points_from_expansions,
- build_v2_user_input,
- )
- from find_agent_v2.gates import build_rule_snapshot, evaluate_candidate_gate
- from find_agent_v2.models import (
- FindAgentV2Candidate,
- FindAgentV2Evidence,
- FindAgentV2Round,
- FindAgentV2Run,
- FindAgentV2Search,
- )
- from find_agent_v2.observability import (
- GRAPH_SPEC,
- MODULE_TITLES,
- OBAGENT_AGENT,
- OBAGENT_PROJECT,
- OBAGENT_ROUND_ANCHOR,
- NullObserver,
- )
- from find_agent_v2.prompts import COMMON_RULES
- from find_agent_v2.providers import normalize_age_pair
- from find_agent_v2.state import DiscoverySnapshot, FindAgentState, NodeRun
- from find_agent_v2.tools import (
- EVALUATION_TOOLS,
- EVIDENCE_TOOLS,
- REPORT_TOOLS,
- SEARCH_TOOLS,
- )
- def _names(functions) -> set[str]:
- return {getattr(fn, "_tool_name", fn.__name__) for fn in functions}
- def test_v2_orm_uses_only_new_table_namespace() -> None:
- assert {
- FindAgentV2Run.__tablename__,
- FindAgentV2Round.__tablename__,
- FindAgentV2Search.__tablename__,
- FindAgentV2Candidate.__tablename__,
- FindAgentV2Evidence.__tablename__,
- } == {
- "find_agent_v2_run",
- "find_agent_v2_round",
- "find_agent_v2_search",
- "find_agent_v2_candidate",
- "find_agent_v2_evidence",
- }
- assert "obagent_run_uid" in FindAgentV2Run.__table__.columns
- assert {"input_tokens", "output_tokens", "total_tokens", "cost_usd"} <= {
- column.name for column in FindAgentV2Run.__table__.columns
- }
- assert not any(table.foreign_key_constraints for table in (
- FindAgentV2Run.__table__,
- FindAgentV2Round.__table__,
- FindAgentV2Search.__table__,
- FindAgentV2Candidate.__table__,
- FindAgentV2Evidence.__table__,
- ))
- def test_v2_package_has_no_legacy_business_imports() -> None:
- package = Path(__file__).parents[2] / "find_agent_v2"
- source = "\n".join(path.read_text(encoding="utf-8") for path in package.glob("*.py"))
- assert "agents.find_agent" not in source
- assert "video_discovery_gates" not in source
- assert "services.video_discovery" not in source
- assert "system_prompt.md" not in source
- def test_project_orm_defines_no_database_foreign_keys() -> None:
- import supply_infra.db.models # noqa: F401
- from supply_infra.db.base import Base
- assert not {
- table.name: sorted(fk.name or "<unnamed>" for fk in table.foreign_key_constraints)
- for table in Base.metadata.tables.values()
- if table.foreign_key_constraints
- }
- def test_v2_owns_age_portrait_normalization() -> None:
- normalized = normalize_age_pair(
- {"年龄": {"50-": {"percentage": "35%", "preference": 120}}},
- {"年龄": {"50岁以上": {"percentage": 0.25, "preference": 110}}},
- )
- assert normalized["content"]["older_ratio"] == 0.35
- assert normalized["account"]["older_ratio"] == 0.25
- assert normalized["consistency"] == "aligned"
- def test_v2_owns_primary_candidate_gate() -> None:
- rules = build_rule_snapshot()
- candidate = {
- "title": "适合父母的实用生活建议",
- "publish_at": rules["current_datetime"],
- "duration_seconds": 60,
- "share_count": 2000,
- "content_50_plus_ratio": 0.35,
- "account_50_plus_ratio": 0.25,
- "relevance_score": 0.8,
- "elder_score": 0.8,
- "share_score": 0.8,
- "value_score": 0.8,
- }
- result = evaluate_candidate_gate(candidate, rules)
- assert result["status"] == "pass"
- assert result["primary_eligible"] is True
- def test_v2_gate_rejects_explicitly_low_evidence() -> None:
- rules = build_rule_snapshot()
- candidate = {
- "title": "普通视频",
- "publish_at": rules["current_datetime"],
- "duration_seconds": 10,
- "share_count": 3,
- "content_50_plus_ratio": 0.01,
- "account_50_plus_ratio": 0.02,
- }
- result = evaluate_candidate_gate(candidate, rules)
- assert result["status"] == "fail"
- assert {"DURATION_TOO_SHORT", "SHARE_COUNT_TOO_LOW", "PORTRAIT_50_PLUS_TOO_LOW"} <= set(
- result["failed_reason_codes"]
- )
- def test_v2_context_deduplicates_expansion_points() -> None:
- rows = [
- SimpleNamespace(
- video_id="v1", point_type="purpose", expanded_text="照顾父母",
- point_desc="描述一",
- ),
- SimpleNamespace(
- video_id="v1", point_type="purpose", expanded_text="照顾父母",
- point_desc="重复描述",
- ),
- SimpleNamespace(
- video_id="v1", point_type="invalid", expanded_text="忽略",
- point_desc=None,
- ),
- ]
- order, points = _points_from_expansions(rows)
- assert order == ["v1"]
- assert len(points["v1"]) == 1
- assert points["v1"][0].point == "照顾父母"
- def test_v2_context_builds_self_contained_user_input() -> None:
- context = V2DemandContext(
- biz_dt="20260811",
- demand_grade_id=123,
- demand_name="测试需求",
- grade="S",
- score=95.0,
- videos=[V2ReferenceVideo(
- video_id="video-1",
- title="参考视频",
- points=[V2ReferencePoint("关键内容", "key", "关键描述")],
- )],
- )
- raw = build_v2_user_input(
- context,
- run_id="v2-test-run",
- rules={"rule_version": "v2-test"},
- )
- assert '"run_id": "v2-test-run"' in raw
- assert '"demand_grade_id": 123' in raw
- assert '"reference_videos"' in raw
- assert '"关键内容"' in raw
- def test_obagent_identity_and_round_structure_are_stable() -> None:
- assert OBAGENT_PROJECT == "find_agent_v2"
- assert OBAGENT_AGENT == "find_agent_v2"
- assert OBAGENT_ROUND_ANCHOR == {"in": "run", "on": ["graph"]}
- assert [node["key"] for node in GRAPH_SPEC["nodes"]] == [
- "supervisor", "search", "evidence", "evaluator",
- ]
- assert set(MODULE_TITLES) == {"supervisor", "search", "evidence", "evaluator", "report"}
- def test_round_graph_is_a_compiled_langgraph() -> None:
- service = _FakeService(pending_after_search=0)
- graph = FindAgentRoundGraph(service=service, runner=_FakeRunner(service))
- drawable = graph.app.get_graph()
- assert {"supervisor", "search", "evidence", "evaluator"} <= set(drawable.nodes)
- assert graph.obagent_spec.get("nodes")
- def test_runtime_exposes_bounded_delegate_schema() -> None:
- field = DelegateArgs.model_fields["requests"]
- assert field.metadata
- assert FindAgentNodeHost.__module__ == "find_agent_v2.runtime"
- def test_runtime_usage_can_be_reset_between_resume_attempts() -> None:
- host = FindAgentNodeHost(observer=NullObserver())
- host.usage["total_tokens"] = 99
- host.reset_usage()
- assert host.usage == {
- "input_tokens": 0,
- "output_tokens": 0,
- "total_tokens": 0,
- "cost": 0.0,
- }
- @pytest.mark.asyncio
- async def test_langchain_runtime_runs_without_network(monkeypatch) -> None:
- host = FindAgentNodeHost(observer=NullObserver())
- fake_model = FakeMessagesListChatModel(responses=[AIMessage(content="done")])
- monkeypatch.setattr(host, "_model", lambda _role: fake_model)
- result = await host.run_node(
- node="supervisor",
- round_index=1,
- system_prompt="plan",
- user_content="task",
- tools=(),
- max_iterations=2,
- allow_delegation=False,
- )
- assert result.content == "done"
- assert result.iterations == 1
- assert result.tool_calls_made == 0
- def test_stage_tool_allowlists_are_physical_and_isolated() -> None:
- assert _names(SEARCH_TOOLS) == {"search_videos_v2", "query_find_agent_v2_state"}
- assert _names(EVIDENCE_TOOLS) == {
- "fetch_candidate_details_v2",
- "fetch_candidate_portraits_v2",
- "query_pending_candidates_v2",
- }
- assert _names(EVALUATION_TOOLS) == {
- "evaluate_candidates_v2",
- "query_pending_candidates_v2",
- }
- assert _names(REPORT_TOOLS) == {"query_find_agent_v2_state"}
- all_names = _names((*SEARCH_TOOLS, *EVIDENCE_TOOLS, *EVALUATION_TOOLS, *REPORT_TOOLS))
- assert not any(name.startswith("batch_update_video_discovery") for name in all_names)
- assert "query_video_discovery_state" not in all_names
- def test_common_prompt_points_to_v2_tables_and_tools() -> None:
- assert "find_agent_v2_run" in COMMON_RULES
- assert "video_discovery_run" not in COMMON_RULES
- assert "batch_search_and_record" not in COMMON_RULES
- assert "batch_update_video_discovery_candidates" not in COMMON_RULES
- class _FakeService:
- def __init__(self, *, pending_after_search: int) -> None:
- self.pending_after_search = pending_after_search
- self.stage = "start"
- self.updates: list[dict] = []
- def get_full_state(self, run_id: str, **_kwargs):
- snapshot = self.snapshot(run_id)
- evidence_status = "pending" if self.stage in {"start", "searched"} else "success"
- return {
- "run": {"run_id": run_id},
- "searches": [],
- "candidates": [
- {
- "candidate_id": index,
- "decision_bucket": "pending_evaluation",
- "detail_status": evidence_status,
- "portrait_status": evidence_status,
- }
- for index in range(1, snapshot.pending_count + 1)
- ],
- }
- def snapshot(self, _run_id: str) -> DiscoverySnapshot:
- base = DiscoverySnapshot("running", 0, 0, 0, 0, 0, 0)
- if self.stage == "searched":
- return replace(base, search_count=1, candidate_count=self.pending_after_search,
- pending_count=self.pending_after_search)
- if self.stage in {"evidenced", "batched"}:
- return replace(base, search_count=1, candidate_count=self.pending_after_search,
- pending_count=self.pending_after_search)
- if self.stage == "evaluated":
- return replace(base, search_count=1, candidate_count=self.pending_after_search,
- rejected_count=self.pending_after_search)
- return base
- def update_round(self, _run_id: str, _round_index: int, **kwargs) -> None:
- self.updates.append(kwargs)
- def recount_valid_primary(self, _run_id: str) -> int:
- return 0
- class _FakeRunner:
- def __init__(self, service: _FakeService) -> None:
- self.service = service
- self.calls: list[tuple[str, set[str]]] = []
- async def run_node(self, *, node, round_index, tools=(), **_kwargs) -> NodeRun:
- self.calls.append((node, _names(tools)))
- if node == "search":
- self.service.stage = "searched"
- elif node == "evidence":
- self.service.stage = "evidenced"
- elif node == "evaluator":
- self.service.stage = "evaluated"
- if node == "supervisor":
- if self.service.stage == "evaluated":
- next_action = "finish"
- elif self.service.stage in {"evidenced", "batched"}:
- next_action = "evaluator"
- elif self.service.stage == "searched" and not self.service.pending_after_search:
- next_action = "finish"
- else:
- next_action = "search"
- content = json.dumps({
- "next_action": next_action,
- "worker_count": 4,
- "evidence_scope": "both",
- "plan": {"searches": []},
- })
- else:
- content = '{"searches": []}'
- return NodeRun(node, round_index, content, 1, 0)
- @pytest.mark.asyncio
- async def test_round_graph_supervisor_routes_with_guarded_allowlists() -> None:
- service = _FakeService(pending_after_search=2)
- runner = _FakeRunner(service)
- graph = FindAgentRoundGraph(service=service, runner=runner)
- state = FindAgentState(run_id="new-run", user_input="task", round_index=1)
- result = await graph.invoke(state)
- assert [name for name, _ in runner.calls] == [
- "supervisor", "search", "supervisor", "evidence", "evidence",
- "supervisor", "evaluator", "supervisor",
- ]
- assert runner.calls[0][1] == set()
- assert runner.calls[1][1] == _names(SEARCH_TOOLS)
- assert runner.calls[2][1] == set()
- assert runner.calls[3][1] <= _names(EVIDENCE_TOOLS)
- assert runner.calls[4][1] <= _names(EVIDENCE_TOOLS)
- assert runner.calls[6][1] == _names(EVALUATION_TOOLS)
- assert result.phase == "done"
- assert result.snapshot is not None and result.snapshot.pending_count == 0
- @pytest.mark.asyncio
- async def test_round_graph_skips_evidence_and_evaluation_without_candidates() -> None:
- service = _FakeService(pending_after_search=0)
- runner = _FakeRunner(service)
- graph = FindAgentRoundGraph(service=service, runner=runner)
- await graph.invoke(FindAgentState(run_id="new-run", user_input="task", round_index=1))
- assert [name for name, _ in runner.calls] == [
- "supervisor", "search", "supervisor",
- ]
- @pytest.mark.asyncio
- async def test_round_graph_reenters_evaluator_until_pending_queue_is_empty() -> None:
- service = _FakeService(pending_after_search=3)
- class BatchedRunner(_FakeRunner):
- def __init__(self, fake_service: _FakeService) -> None:
- super().__init__(fake_service)
- self.remaining = 3
- async def run_node(self, *, node, round_index, tools=(), **kwargs) -> NodeRun:
- result = await super().run_node(
- node=node, round_index=round_index, tools=tools, **kwargs,
- )
- if node == "evaluator":
- self.remaining -= 1
- self.service.stage = "evaluated" if self.remaining == 0 else "batched"
- return result
- runner = BatchedRunner(service)
- original_snapshot = service.snapshot
- def batched_snapshot(run_id: str) -> DiscoverySnapshot:
- if service.stage == "batched":
- base = original_snapshot(run_id)
- return replace(
- base,
- search_count=1,
- candidate_count=3,
- pending_count=runner.remaining,
- rejected_count=3 - runner.remaining,
- )
- return original_snapshot(run_id)
- service.snapshot = batched_snapshot # type: ignore[method-assign]
- graph = FindAgentRoundGraph(service=service, runner=runner)
- result = await graph.invoke(
- FindAgentState(run_id="batched-run", user_input="task", round_index=1),
- )
- assert [name for name, _ in runner.calls].count("evaluator") == 3
- assert result.snapshot is not None and result.snapshot.pending_count == 0
- def test_supervisor_policy_overrides_unsafe_finish_and_clamps_workers() -> None:
- service = _FakeService(pending_after_search=2)
- service.stage = "searched"
- graph = FindAgentRoundGraph(service=service, runner=_FakeRunner(service))
- action, _reason, workers, scope = graph._approve_action(
- {"run_id": "r", "search_actions": 1, "action_count": 1},
- {"next_action": "finish", "worker_count": 99, "evidence_scope": "invalid"},
- )
- assert action == "evidence"
- assert workers == 8
- assert scope == "both"
- def test_supervisor_can_choose_an_extra_search_within_budget() -> None:
- service = _FakeService(pending_after_search=0)
- service.stage = "searched"
- graph = FindAgentRoundGraph(service=service, runner=_FakeRunner(service))
- action, *_ = graph._approve_action(
- {"run_id": "r", "search_actions": 1, "action_count": 1},
- {"next_action": "search", "worker_count": 2},
- )
- assert action == "search"
- @pytest.mark.asyncio
- async def test_round_graph_retries_one_stagnant_evaluator_response() -> None:
- service = _FakeService(pending_after_search=2)
- class RetryRunner(_FakeRunner):
- evaluator_calls = 0
- async def run_node(self, *, node, round_index, tools=(), **kwargs) -> NodeRun:
- if node == "evaluator":
- self.calls.append((node, _names(tools)))
- self.evaluator_calls += 1
- if self.evaluator_calls == 2:
- self.service.stage = "evaluated"
- return NodeRun(node, round_index, "", 1, 0)
- return await super().run_node(
- node=node, round_index=round_index, tools=tools, **kwargs,
- )
- runner = RetryRunner(service)
- graph = FindAgentRoundGraph(service=service, runner=runner)
- result = await graph.invoke(
- FindAgentState(run_id="retry-run", user_input="task", round_index=1),
- )
- assert runner.evaluator_calls == 2
- assert result.snapshot is not None and result.snapshot.pending_count == 0
|