graph.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302
  1. from __future__ import annotations
  2. from concurrent.futures import ThreadPoolExecutor
  3. from dataclasses import dataclass
  4. from datetime import datetime, timezone
  5. from typing import Literal
  6. from langgraph.graph import END, START, StateGraph
  7. from app.config import Settings
  8. from app.decision_engine import DecisionEngine
  9. from app.schemas import RouteName, ToolResult
  10. from app.state import AgentState
  11. from app.tools import SQLQueryTool, VectorSearchTool, WebSearchTool
  12. def trace_event(state: AgentState, node: str, detail: dict) -> list[dict]:
  13. return [
  14. *state.get("trace", []),
  15. {
  16. "node": node,
  17. "at": datetime.now(timezone.utc).isoformat(),
  18. "detail": detail,
  19. },
  20. ]
  21. @dataclass
  22. class GraphDependencies:
  23. settings: Settings
  24. engine: DecisionEngine
  25. vector_tool: VectorSearchTool
  26. sql_tool: SQLQueryTool
  27. web_tool: WebSearchTool
  28. def build_graph(deps: GraphDependencies):
  29. def normalize_query(state: AgentState) -> dict:
  30. normalized = " ".join(state["original_query"].strip().split())
  31. return {
  32. "current_query": normalized,
  33. "retrieval_round": state.get("retrieval_round", 0),
  34. "trace": trace_event(state, "normalize_query", {"query": normalized}),
  35. }
  36. def route_query(state: AgentState) -> dict:
  37. decision = deps.engine.route(
  38. state["current_query"], deps.settings.max_retrieval_rounds
  39. )
  40. return {
  41. "route_decision": decision,
  42. "trace": trace_event(
  43. state,
  44. "route_query",
  45. {
  46. "intent": decision.intent,
  47. "routes": [route.value for route in decision.routes],
  48. "reason_code": decision.reason_code,
  49. },
  50. ),
  51. }
  52. def after_route(
  53. state: AgentState,
  54. ) -> Literal["clarify", "refuse", "generate", "plan"]:
  55. routes = set(state["route_decision"].routes)
  56. # 退出类路由优先级高于检索,防止复合输出意外触发工具。
  57. if RouteName.CLARIFY in routes:
  58. return "clarify"
  59. if RouteName.REFUSE in routes:
  60. return "refuse"
  61. if RouteName.DIRECT_ANSWER in routes:
  62. return "generate"
  63. return "plan"
  64. def plan_query(state: AgentState) -> dict:
  65. plan = deps.engine.plan(state["current_query"], state["route_decision"])
  66. return {
  67. "retrieval_plan": plan,
  68. "trace": trace_event(
  69. state,
  70. "plan_query",
  71. {"steps": [step.model_dump(mode="json") for step in plan.steps]},
  72. ),
  73. }
  74. def execute_step(step) -> ToolResult:
  75. # Executor 只识别注册过的枚举工具;模型不能构造任意函数名。
  76. if step.tool == RouteName.MILVUS_SEARCH:
  77. return deps.vector_tool.invoke(
  78. query=step.query,
  79. policy_type=str(step.arguments.get("policy_type", "")),
  80. top_k=int(step.arguments.get("top_k", 5)),
  81. )
  82. if step.tool == RouteName.SQL_QUERY:
  83. return deps.sql_tool.invoke(
  84. user_id=str(step.arguments.get("user_id", "")),
  85. days=int(step.arguments.get("days", 30)),
  86. product_keyword=str(step.arguments.get("product_keyword", "")),
  87. )
  88. if step.tool == RouteName.WEB_SEARCH:
  89. requested = int(step.arguments.get("max_results", 5))
  90. # 即使模型给出更大的 max_results,也不能突破程序配置上限。
  91. return deps.web_tool.invoke(
  92. query=step.query,
  93. max_results=min(max(requested, 1), deps.settings.web_search_max_results),
  94. )
  95. return ToolResult(
  96. status="error",
  97. tool=step.tool.value,
  98. error_code="UNSUPPORTED_TOOL",
  99. error_message=f"未注册工具:{step.tool.value}",
  100. )
  101. def execute_plan(state: AgentState) -> dict:
  102. steps = state["retrieval_plan"].steps
  103. if not steps:
  104. results: list[ToolResult] = []
  105. elif len(steps) == 1:
  106. results = [execute_step(steps[0])]
  107. else:
  108. # 本项目的计划均为无依赖读任务;依赖型 DAG 需按拓扑批次调度。
  109. with ThreadPoolExecutor(max_workers=min(len(steps), 4)) as executor:
  110. results = list(executor.map(execute_step, steps))
  111. # ToolResult 用于控制流,Evidence 用于答案生成,二者不能混为一体。
  112. evidence = [item for result in results for item in result.evidence]
  113. errors = [
  114. {
  115. "tool": result.tool,
  116. "code": result.error_code,
  117. "message": result.error_message,
  118. }
  119. for result in results
  120. if result.status == "error"
  121. ]
  122. # 将规划 Query、受控参数和实际结果并列记录,便于观察真实调用。
  123. executed_queries = [
  124. {
  125. "step_id": step.id,
  126. "tool": result.tool,
  127. "query": step.query,
  128. "arguments": step.arguments,
  129. "status": result.status,
  130. "latency_ms": result.latency_ms,
  131. }
  132. for step, result in zip(steps, results, strict=True)
  133. ]
  134. return {
  135. "tool_results": results,
  136. "executed_queries": executed_queries,
  137. "evidence": evidence,
  138. "errors": [*state.get("errors", []), *errors],
  139. "trace": trace_event(
  140. state,
  141. "execute_plan",
  142. {
  143. "tools": [result.tool for result in results],
  144. "statuses": [result.status for result in results],
  145. "evidence_count": len(evidence),
  146. "executed_queries": executed_queries,
  147. },
  148. ),
  149. }
  150. def grade_evidence(state: AgentState) -> dict:
  151. current_round = state.get("retrieval_round", 0)
  152. max_rounds = state["route_decision"].max_rounds
  153. grade = deps.engine.grade(
  154. query=state["current_query"],
  155. decision=state["route_decision"],
  156. evidence=state.get("evidence", []),
  157. current_round=current_round,
  158. max_rounds=max_rounds,
  159. min_score=deps.settings.min_evidence_score,
  160. )
  161. # 模型只能建议 Rewrite,是否还有预算必须由程序层决定。
  162. # 达到上限后强制 Stop,避免模型持续 Rewrite 导致图无限循环。
  163. if (
  164. grade.recommended_action == "rewrite_query"
  165. and current_round >= max_rounds
  166. ):
  167. grade = grade.model_copy(
  168. update={
  169. "sufficient": False,
  170. "recommended_action": "stop",
  171. "reason": (
  172. f"{grade.reason};已达到最大检索轮数 {max_rounds},"
  173. "程序层强制停止。"
  174. ),
  175. }
  176. )
  177. return {
  178. "quality_grade": grade,
  179. "trace": trace_event(
  180. state,
  181. "grade_evidence",
  182. grade.model_dump(mode="json"),
  183. ),
  184. }
  185. def after_grade(state: AgentState) -> Literal["generate", "rewrite", "generate_partial"]:
  186. # 条件边只消费结构化动作,不解析 grader 的自然语言 reason。
  187. action = state["quality_grade"].recommended_action
  188. if action == "accept":
  189. return "generate"
  190. if action == "rewrite_query":
  191. return "rewrite"
  192. return "generate_partial"
  193. def rewrite_query(state: AgentState) -> dict:
  194. rewritten = deps.engine.rewrite(
  195. state["current_query"], state["quality_grade"]
  196. )
  197. next_round = state.get("retrieval_round", 0) + 1
  198. return {
  199. "current_query": rewritten,
  200. "retrieval_round": next_round,
  201. "trace": trace_event(
  202. state,
  203. "rewrite_query",
  204. {"round": next_round, "rewritten_query": rewritten},
  205. ),
  206. }
  207. def generate_answer(state: AgentState) -> dict:
  208. answer = deps.engine.answer(
  209. state["original_query"], state.get("evidence", []), partial=False
  210. )
  211. return {
  212. "final_answer": answer,
  213. "termination_reason": "evidence_accepted"
  214. if state.get("evidence")
  215. else "direct_answer",
  216. "trace": trace_event(state, "generate_answer", {"partial": False}),
  217. }
  218. def generate_partial_answer(state: AgentState) -> dict:
  219. answer = deps.engine.answer(
  220. state["original_query"], state.get("evidence", []), partial=True
  221. )
  222. return {
  223. "final_answer": answer,
  224. "termination_reason": "retrieval_budget_exhausted",
  225. "trace": trace_event(state, "generate_partial_answer", {"partial": True}),
  226. }
  227. def clarify(state: AgentState) -> dict:
  228. return {
  229. "final_answer": "请补充用户编号、商品名称或订单范围后再查询。",
  230. "termination_reason": "clarification_required",
  231. "trace": trace_event(state, "clarify", {}),
  232. }
  233. def refuse(state: AgentState) -> dict:
  234. return {
  235. "final_answer": "当前请求涉及受限数据,系统拒绝执行。",
  236. "termination_reason": "security_policy",
  237. "trace": trace_event(state, "refuse", {}),
  238. }
  239. builder = StateGraph(AgentState)
  240. builder.add_node("normalize_query", normalize_query)
  241. builder.add_node("route_query", route_query)
  242. builder.add_node("plan_query", plan_query)
  243. builder.add_node("execute_plan", execute_plan)
  244. builder.add_node("grade_evidence", grade_evidence)
  245. builder.add_node("rewrite_query", rewrite_query)
  246. builder.add_node("generate_answer", generate_answer)
  247. builder.add_node("generate_partial_answer", generate_partial_answer)
  248. builder.add_node("clarify", clarify)
  249. builder.add_node("refuse", refuse)
  250. builder.add_edge(START, "normalize_query")
  251. builder.add_edge("normalize_query", "route_query")
  252. builder.add_conditional_edges(
  253. "route_query",
  254. after_route,
  255. {
  256. "clarify": "clarify",
  257. "refuse": "refuse",
  258. "generate": "generate_answer",
  259. "plan": "plan_query",
  260. },
  261. )
  262. builder.add_edge("plan_query", "execute_plan")
  263. builder.add_edge("execute_plan", "grade_evidence")
  264. builder.add_conditional_edges(
  265. "grade_evidence",
  266. after_grade,
  267. {
  268. "generate": "generate_answer",
  269. "rewrite": "rewrite_query",
  270. "generate_partial": "generate_partial_answer",
  271. },
  272. )
  273. builder.add_edge("rewrite_query", "plan_query")
  274. builder.add_edge("generate_answer", END)
  275. builder.add_edge("generate_partial_answer", END)
  276. builder.add_edge("clarify", END)
  277. builder.add_edge("refuse", END)
  278. return builder.compile()