from pathlib import Path from app.config import Settings from app.decision_engine import DemoDecisionEngine from app.schemas import Evidence, QualityGrade, QueryRequest, RouteName from app.service import AgenticRAGService from app.sql_store import initialize_database class FakeVectorStore: def __init__(self, return_evidence: bool = True) -> None: self.return_evidence = return_evidence self.calls = 0 def search( self, query: str, top_k: int = 5, policy_type: str = "", ) -> list[Evidence]: self.calls += 1 if not self.return_evidence: return [] return [ Evidence( source_type="milvus", source="shipping_policy.md", content="单笔有效订单实付满 99 元免基础运费。", score=0.91, metadata={"policy_type": policy_type}, ) ] class FakeWebSearchProvider: def __init__(self) -> None: self.calls = 0 def search(self, query: str, max_results: int = 5) -> list[Evidence]: self.calls += 1 return [ Evidence( source_type="web", source="https://brand.example/xphone-15-pro-notice", content="品牌官网发布了 XPhone 15 Pro 最新服务公告。", score=0.88, metadata={"title": "XPhone 15 Pro 最新公告"}, ) ] class AlwaysRewriteDecisionEngine(DemoDecisionEngine): """模拟模型忽略检索预算、每轮都要求继续改写。""" def grade(self, *args, **kwargs) -> QualityGrade: return QualityGrade( relevant=False, sufficient=False, missing_aspects=["网页证据"], recommended_action="rewrite_query", reason="模拟模型持续要求 Rewrite", ) def create_settings(sqlite_path: Path, max_rounds: int = 2) -> Settings: return Settings( llm_provider="demo", sqlite_path=sqlite_path, max_retrieval_rounds=max_rounds, min_evidence_score=0.35, ) def test_multi_source_graph(tmp_path) -> None: database = tmp_path / "shop.db" initialize_database(database, reset=True) service = AgenticRAGService(create_settings(database), FakeVectorStore()) response = service.invoke( QueryRequest( query="统计 U1001 最近 30 天订单,并结合内部运费规则说明是否包邮", debug=True, ) ) assert response.route.routes == [RouteName.MILVUS_SEARCH, RouteName.SQL_QUERY] assert response.termination_reason == "evidence_accepted" assert {item.source_type for item in response.citations} == {"milvus", "sql"} assert {item["tool"] for item in response.executed_queries} == { "milvus_search", "sql_query", } assert all(item["query"] for item in response.executed_queries) assert [event["node"] for event in response.trace] == [ "normalize_query", "route_query", "plan_query", "execute_plan", "grade_evidence", "generate_answer", ] def test_empty_retrieval_stops_after_budget(tmp_path) -> None: database = tmp_path / "shop.db" initialize_database(database, reset=True) vector_store = FakeVectorStore(return_evidence=False) service = AgenticRAGService( create_settings(database, max_rounds=1), vector_store ) response = service.invoke( QueryRequest(query="耳机拆封后还能七日无理由退货吗?", debug=True) ) assert vector_store.calls == 2 assert response.termination_reason == "retrieval_budget_exhausted" assert "未获得" in response.answer assert [event["node"] for event in response.trace].count("rewrite_query") == 1 def test_graph_forces_stop_when_model_ignores_rewrite_budget( tmp_path, monkeypatch ) -> None: database = tmp_path / "shop.db" initialize_database(database, reset=True) vector_store = FakeVectorStore(return_evidence=False) monkeypatch.setattr( "app.service.create_decision_engine", lambda settings: AlwaysRewriteDecisionEngine(), ) service = AgenticRAGService( create_settings(database, max_rounds=1), vector_store, ) response = service.invoke( QueryRequest(query="耳机退货政策", debug=True) ) assert vector_store.calls == 2 assert response.termination_reason == "retrieval_budget_exhausted" assert response.trace[-2]["node"] == "grade_evidence" assert response.trace[-2]["detail"]["recommended_action"] == "stop" def test_three_source_graph(tmp_path) -> None: database = tmp_path / "shop.db" initialize_database(database, reset=True) web_provider = FakeWebSearchProvider() service = AgenticRAGService( create_settings(database), FakeVectorStore(), web_provider, ) response = service.invoke( QueryRequest( query=( "统计 U1001 最近 30 天购买 XPhone 15 Pro 的订单金额," "结合内部退货政策和品牌官网最新公告给出售后建议" ), debug=True, ) ) assert response.route.routes == [ RouteName.MILVUS_SEARCH, RouteName.SQL_QUERY, RouteName.WEB_SEARCH, ] assert web_provider.calls == 1 assert {item.source_type for item in response.citations} == { "milvus", "sql", "web", } assert response.termination_reason == "evidence_accepted" assert len(response.executed_queries) == 3 sql_call = next( item for item in response.executed_queries if item["tool"] == "sql_query" ) assert sql_call["arguments"]["user_id"] == "U1001" assert sql_call["arguments"]["product_keyword"] == "XPhone 15 Pro" def test_clarification_does_not_call_vector_store(tmp_path) -> None: database = tmp_path / "shop.db" initialize_database(database, reset=True) vector_store = FakeVectorStore() service = AgenticRAGService(create_settings(database), vector_store) response = service.invoke(QueryRequest(query="我最近 30 天买了几单?")) assert response.termination_reason == "clarification_required" assert vector_store.calls == 0