test_graph.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196
  1. from pathlib import Path
  2. from app.config import Settings
  3. from app.decision_engine import DemoDecisionEngine
  4. from app.schemas import Evidence, QualityGrade, QueryRequest, RouteName
  5. from app.service import AgenticRAGService
  6. from app.sql_store import initialize_database
  7. class FakeVectorStore:
  8. def __init__(self, return_evidence: bool = True) -> None:
  9. self.return_evidence = return_evidence
  10. self.calls = 0
  11. def search(
  12. self,
  13. query: str,
  14. top_k: int = 5,
  15. policy_type: str = "",
  16. ) -> list[Evidence]:
  17. self.calls += 1
  18. if not self.return_evidence:
  19. return []
  20. return [
  21. Evidence(
  22. source_type="milvus",
  23. source="shipping_policy.md",
  24. content="单笔有效订单实付满 99 元免基础运费。",
  25. score=0.91,
  26. metadata={"policy_type": policy_type},
  27. )
  28. ]
  29. class FakeWebSearchProvider:
  30. def __init__(self) -> None:
  31. self.calls = 0
  32. def search(self, query: str, max_results: int = 5) -> list[Evidence]:
  33. self.calls += 1
  34. return [
  35. Evidence(
  36. source_type="web",
  37. source="https://brand.example/xphone-15-pro-notice",
  38. content="品牌官网发布了 XPhone 15 Pro 最新服务公告。",
  39. score=0.88,
  40. metadata={"title": "XPhone 15 Pro 最新公告"},
  41. )
  42. ]
  43. class AlwaysRewriteDecisionEngine(DemoDecisionEngine):
  44. """模拟模型忽略检索预算、每轮都要求继续改写。"""
  45. def grade(self, *args, **kwargs) -> QualityGrade:
  46. return QualityGrade(
  47. relevant=False,
  48. sufficient=False,
  49. missing_aspects=["网页证据"],
  50. recommended_action="rewrite_query",
  51. reason="模拟模型持续要求 Rewrite",
  52. )
  53. def create_settings(sqlite_path: Path, max_rounds: int = 2) -> Settings:
  54. return Settings(
  55. llm_provider="demo",
  56. sqlite_path=sqlite_path,
  57. max_retrieval_rounds=max_rounds,
  58. min_evidence_score=0.35,
  59. )
  60. def test_multi_source_graph(tmp_path) -> None:
  61. database = tmp_path / "shop.db"
  62. initialize_database(database, reset=True)
  63. service = AgenticRAGService(create_settings(database), FakeVectorStore())
  64. response = service.invoke(
  65. QueryRequest(
  66. query="统计 U1001 最近 30 天订单,并结合内部运费规则说明是否包邮",
  67. debug=True,
  68. )
  69. )
  70. assert response.route.routes == [RouteName.MILVUS_SEARCH, RouteName.SQL_QUERY]
  71. assert response.termination_reason == "evidence_accepted"
  72. assert {item.source_type for item in response.citations} == {"milvus", "sql"}
  73. assert {item["tool"] for item in response.executed_queries} == {
  74. "milvus_search",
  75. "sql_query",
  76. }
  77. assert all(item["query"] for item in response.executed_queries)
  78. assert [event["node"] for event in response.trace] == [
  79. "normalize_query",
  80. "route_query",
  81. "plan_query",
  82. "execute_plan",
  83. "grade_evidence",
  84. "generate_answer",
  85. ]
  86. def test_empty_retrieval_stops_after_budget(tmp_path) -> None:
  87. database = tmp_path / "shop.db"
  88. initialize_database(database, reset=True)
  89. vector_store = FakeVectorStore(return_evidence=False)
  90. service = AgenticRAGService(
  91. create_settings(database, max_rounds=1), vector_store
  92. )
  93. response = service.invoke(
  94. QueryRequest(query="耳机拆封后还能七日无理由退货吗?", debug=True)
  95. )
  96. assert vector_store.calls == 2
  97. assert response.termination_reason == "retrieval_budget_exhausted"
  98. assert "未获得" in response.answer
  99. assert [event["node"] for event in response.trace].count("rewrite_query") == 1
  100. def test_graph_forces_stop_when_model_ignores_rewrite_budget(
  101. tmp_path, monkeypatch
  102. ) -> None:
  103. database = tmp_path / "shop.db"
  104. initialize_database(database, reset=True)
  105. vector_store = FakeVectorStore(return_evidence=False)
  106. monkeypatch.setattr(
  107. "app.service.create_decision_engine",
  108. lambda settings: AlwaysRewriteDecisionEngine(),
  109. )
  110. service = AgenticRAGService(
  111. create_settings(database, max_rounds=1),
  112. vector_store,
  113. )
  114. response = service.invoke(
  115. QueryRequest(query="耳机退货政策", debug=True)
  116. )
  117. assert vector_store.calls == 2
  118. assert response.termination_reason == "retrieval_budget_exhausted"
  119. assert response.trace[-2]["node"] == "grade_evidence"
  120. assert response.trace[-2]["detail"]["recommended_action"] == "stop"
  121. def test_three_source_graph(tmp_path) -> None:
  122. database = tmp_path / "shop.db"
  123. initialize_database(database, reset=True)
  124. web_provider = FakeWebSearchProvider()
  125. service = AgenticRAGService(
  126. create_settings(database),
  127. FakeVectorStore(),
  128. web_provider,
  129. )
  130. response = service.invoke(
  131. QueryRequest(
  132. query=(
  133. "统计 U1001 最近 30 天购买 XPhone 15 Pro 的订单金额,"
  134. "结合内部退货政策和品牌官网最新公告给出售后建议"
  135. ),
  136. debug=True,
  137. )
  138. )
  139. assert response.route.routes == [
  140. RouteName.MILVUS_SEARCH,
  141. RouteName.SQL_QUERY,
  142. RouteName.WEB_SEARCH,
  143. ]
  144. assert web_provider.calls == 1
  145. assert {item.source_type for item in response.citations} == {
  146. "milvus",
  147. "sql",
  148. "web",
  149. }
  150. assert response.termination_reason == "evidence_accepted"
  151. assert len(response.executed_queries) == 3
  152. sql_call = next(
  153. item for item in response.executed_queries if item["tool"] == "sql_query"
  154. )
  155. assert sql_call["arguments"]["user_id"] == "U1001"
  156. assert sql_call["arguments"]["product_keyword"] == "XPhone 15 Pro"
  157. def test_clarification_does_not_call_vector_store(tmp_path) -> None:
  158. database = tmp_path / "shop.db"
  159. initialize_database(database, reset=True)
  160. vector_store = FakeVectorStore()
  161. service = AgenticRAGService(create_settings(database), vector_store)
  162. response = service.invoke(QueryRequest(query="我最近 30 天买了几单?"))
  163. assert response.termination_reason == "clarification_required"
  164. assert vector_store.calls == 0