graph.py 9.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291
  1. from app.tools import SQLQueryTool,WebSearchTool,VectorSearchTool
  2. from app.decision_engine import DeepSeekDecisionEngine
  3. from app.config import Settings
  4. from dataclasses import dataclass
  5. from concurrent.futures import ThreadPoolExecutor
  6. from langgraph.graph import END, START, StateGraph
  7. from app.state import AgentState
  8. from app.schemas import RouteName, ToolResult,PlanStep
  9. @dataclass
  10. class GraphBuilder:
  11. settings:Settings
  12. engine:DeepSeekDecisionEngine
  13. sql_tool:SQLQueryTool
  14. vector_tool:VectorSearchTool
  15. web_tool:WebSearchTool
  16. def trace_event(state:AgentState,node:str,detail:dict)->list[dict]:
  17. return[
  18. *state.get("trace",[]),
  19. {
  20. "node":node,
  21. "detail":detail
  22. }
  23. ]
  24. def build_graph(builder:GraphBuilder):
  25. def normalize_query(state: AgentState) -> dict:
  26. normalized = " ".join(state["original_query"].strip().split())
  27. return {
  28. "current_query": normalized,
  29. "retrieval_round": state.get("retrieval_round", 0),
  30. "trace": trace_event(state, "normalize_query", {"query": normalized}),
  31. }
  32. def route_query(state:AgentState)->dict:
  33. decision=builder.engine.route(state["current_query"],builder.settings.max_retrieval_rounds)
  34. return{
  35. "route_decision":decision,
  36. "trace":trace_event(
  37. state,
  38. node="route_query",
  39. detail={
  40. "routes": decision.routes
  41. }
  42. )
  43. }
  44. def after_route(state:AgentState)->str:
  45. routes=set(state["route_decision"].routes)
  46. if RouteName.CLARIFY in routes:
  47. return "clarify"
  48. if RouteName.REFUSE in routes:
  49. return "refuse"
  50. if RouteName.DIRECT_ANSWER in routes:
  51. return "generate"
  52. return "plan"
  53. def plan_query(state:AgentState)->dict:
  54. plan=builder.engine.plan(state["current_query"],state["route_decision"])
  55. return{
  56. "retrieval_plan":plan,
  57. "trace":trace_event(
  58. state,
  59. node="plan_query",
  60. detail={
  61. "steps":plan.steps
  62. }
  63. )
  64. }
  65. def execute_step(step:PlanStep)->ToolResult:
  66. if step.tool==RouteName.MILVUS_SEARCH:
  67. return builder.vector_tool.invoke(
  68. query=step.query,
  69. service=step.arguments.get("service",""),
  70. fault_type=step.arguments.get("fault_type",""),
  71. client_os=step.arguments.get("client_os",""),
  72. top_k=step.arguments.get("top_k",5)
  73. )
  74. if step.tool==RouteName.SQL_QUERY:
  75. return builder.sql_tool.invoke(
  76. service=step.arguments.get("service",""),
  77. fault_type=step.arguments.get("fault_type",""),
  78. days=step.arguments.get("days",30),
  79. client_os=step.arguments.get("client_os","")
  80. )
  81. if step.tool==RouteName.WEB_SEARCH:
  82. return builder.web_tool.invoke(
  83. query=step.query,
  84. max_results=step.arguments.get("max_results",5)
  85. )
  86. return ToolResult(
  87. status="error",
  88. tool=step.tool.value,
  89. error_code="UNSUPPORTED_TOOL",
  90. error_message=f"未注册工具:{step.tool.value}",
  91. )
  92. def execute_plan(state:AgentState)->dict:
  93. steps=state["retrieval_plan"].steps
  94. if not steps:
  95. results:list[ToolResult]=[]
  96. else:
  97. with ThreadPoolExecutor(max_workers=min(len(steps),4)) as executor:
  98. results=list(executor.map(execute_step, steps))
  99. evidence=[]
  100. for result in results:
  101. evidence.extend(result.evidence)
  102. errors = [
  103. {
  104. "tool": result.tool,
  105. "code": result.error_code,
  106. "message": result.error_message,
  107. }
  108. for result in results
  109. if result.status == "error"
  110. ]
  111. executed_queries = [
  112. {
  113. "step_id": step.id,
  114. "tool": result.tool,
  115. "query": step.query,
  116. "arguments": step.arguments,
  117. "status": result.status,
  118. "latency_ms": result.latency_ms,
  119. }
  120. for step, result in zip(steps, results, strict=True)
  121. ]
  122. return {
  123. "tool_results": results,
  124. "executed_queries": executed_queries,
  125. "evidence": evidence,
  126. "errors": [*state.get("errors", []), *errors],
  127. "trace": trace_event(
  128. state,
  129. "execute_plan",
  130. {
  131. "tools": [result.tool for result in results],
  132. "statuses": [result.status for result in results],
  133. "evidence_count": len(evidence),
  134. "executed_queries": executed_queries,
  135. },
  136. ),
  137. }
  138. def grade_evidence(state:AgentState)->dict:
  139. current_round = state.get("retrieval_round", 0)
  140. max_rounds = state["route_decision"].max_rounds
  141. grade=builder.engine.grade(
  142. query=state["current_query"],
  143. decision=state["route_decision"],
  144. evidence=state["evidence"],
  145. current_round=current_round,
  146. max_rounds=max_rounds,
  147. min_score=builder.settings.min_evidence_score
  148. )
  149. if(grade.recommended_action=="rewrite_query"
  150. and current_round >= max_rounds):
  151. grade = grade.model_copy(
  152. update={
  153. "sufficient": False,
  154. "recommended_action": "stop",
  155. "reason": (
  156. f"{grade.reason};已达到最大检索轮数 {max_rounds},"
  157. "程序层强制停止。"
  158. ),
  159. }
  160. )
  161. return {
  162. "quality_grade":grade,
  163. "trace":trace_event(
  164. state=state,
  165. node="grade_evidence",
  166. detail=grade.model_dump(mode="json"),
  167. ),
  168. }
  169. def after_grade(state:AgentState)->str:
  170. action=state["quality_grade"].recommended_action
  171. if action=="accept":
  172. return "generate"
  173. if action=="rewrite_query":
  174. return "rewrite"
  175. return "generate_partial"
  176. def rewrite_query(state:AgentState)->dict:
  177. rewritten=builder.engine.rewrite(
  178. query=state["current_query"],
  179. grade=state["quality_grade"])
  180. next_round=state.get("retrieval_round",0)+1
  181. return{
  182. "current_query":rewritten,
  183. "retrieval_round":next_round,
  184. "trace":trace_event(state=state,
  185. node="rewrite_query",
  186. detail={
  187. "round":next_round,
  188. "rewritten":rewritten
  189. })
  190. }
  191. def generate_answer(state: AgentState) -> dict:
  192. answer = builder.engine.answer(
  193. state["original_query"], state.get("evidence", []), partial=False
  194. )
  195. return {
  196. "final_answer": answer,
  197. "termination_reason": "evidence_accepted"
  198. if state.get("evidence")
  199. else "direct_answer",
  200. "trace": trace_event(state, "generate_answer", {"partial": False}),
  201. }
  202. def generate_partial_answer(state: AgentState) -> dict:
  203. answer = builder.engine.answer(
  204. state["original_query"], state.get("evidence", []), partial=True
  205. )
  206. return {
  207. "final_answer": answer,
  208. "termination_reason": "retrieval_budget_exhausted",
  209. "trace": trace_event(state, "generate_partial_answer", {"partial": True}),
  210. }
  211. def clarify(state: AgentState) -> dict:
  212. return {
  213. "final_answer": "当前流程仅处理阿里云 SSL-VPN 的认证超时故障;请确认是否出现认证超时提示,并补充客户端系统或版本信息。",
  214. "termination_reason": "clarification_required",
  215. "trace": trace_event(state, "clarify", {}),
  216. }
  217. def refuse(state: AgentState) -> dict:
  218. return {
  219. "final_answer": "当前请求涉及受限数据,系统拒绝执行。",
  220. "termination_reason": "security_policy",
  221. "trace": trace_event(state, "refuse", {}),
  222. }
  223. graph=StateGraph(AgentState)
  224. graph.add_node("normalize_query",normalize_query)
  225. graph.add_node("route_query",route_query)
  226. graph.add_node("plan_query",plan_query)
  227. graph.add_node("execute_plan",execute_plan)
  228. graph.add_node("grade_evidence",grade_evidence)
  229. graph.add_node("rewrite_query",rewrite_query)
  230. graph.add_node("generate_answer",generate_answer)
  231. graph.add_node("generate_partial_answer",generate_partial_answer)
  232. graph.add_node("clarify",clarify)
  233. graph.add_node("refuse",refuse)
  234. graph.add_edge(START,"normalize_query")
  235. graph.add_edge("normalize_query","route_query")
  236. graph.add_conditional_edges(
  237. "route_query",
  238. after_route,
  239. {
  240. "clarify": "clarify",
  241. "refuse": "refuse",
  242. "generate": "generate_answer",
  243. "plan": "plan_query",
  244. }
  245. )
  246. graph.add_edge("plan_query","execute_plan")
  247. graph.add_edge("execute_plan","grade_evidence")
  248. graph.add_conditional_edges(
  249. "grade_evidence",
  250. after_grade,
  251. {
  252. "generate":"generate_answer",
  253. "rewrite":"rewrite_query",
  254. "generate_partial":"generate_partial_answer"
  255. }
  256. )
  257. graph.add_edge("rewrite_query","plan_query")
  258. graph.add_edge("generate_answer",END)
  259. graph.add_edge("generate_partial_answer",END)
  260. graph.add_edge("clarify", END)
  261. graph.add_edge("refuse", END)
  262. return graph.compile()