| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196 |
- 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
|