|
|
@@ -0,0 +1,291 @@
|
|
|
+from app.tools import SQLQueryTool,WebSearchTool,VectorSearchTool
|
|
|
+from app.decision_engine import DeepSeekDecisionEngine
|
|
|
+from app.config import Settings
|
|
|
+from dataclasses import dataclass
|
|
|
+
|
|
|
+from concurrent.futures import ThreadPoolExecutor
|
|
|
+from langgraph.graph import END, START, StateGraph
|
|
|
+from app.state import AgentState
|
|
|
+from app.schemas import RouteName, ToolResult,PlanStep
|
|
|
+
|
|
|
+@dataclass
|
|
|
+class GraphBuilder:
|
|
|
+ settings:Settings
|
|
|
+ engine:DeepSeekDecisionEngine
|
|
|
+ sql_tool:SQLQueryTool
|
|
|
+ vector_tool:VectorSearchTool
|
|
|
+ web_tool:WebSearchTool
|
|
|
+
|
|
|
+def trace_event(state:AgentState,node:str,detail:dict)->list[dict]:
|
|
|
+ return[
|
|
|
+ *state.get("trace",[]),
|
|
|
+ {
|
|
|
+ "node":node,
|
|
|
+ "detail":detail
|
|
|
+ }
|
|
|
+ ]
|
|
|
+
|
|
|
+def build_graph(builder:GraphBuilder):
|
|
|
+ def normalize_query(state: AgentState) -> dict:
|
|
|
+ normalized = " ".join(state["original_query"].strip().split())
|
|
|
+ return {
|
|
|
+ "current_query": normalized,
|
|
|
+ "retrieval_round": state.get("retrieval_round", 0),
|
|
|
+ "trace": trace_event(state, "normalize_query", {"query": normalized}),
|
|
|
+ }
|
|
|
+
|
|
|
+ def route_query(state:AgentState)->dict:
|
|
|
+ decision=builder.engine.route(state["current_query"],builder.settings.max_retrieval_rounds)
|
|
|
+ return{
|
|
|
+ "route_decision":decision,
|
|
|
+ "trace":trace_event(
|
|
|
+ state,
|
|
|
+ node="route_query",
|
|
|
+ detail={
|
|
|
+ "routes": decision.routes
|
|
|
+ }
|
|
|
+ )
|
|
|
+ }
|
|
|
+
|
|
|
+ def after_route(state:AgentState)->str:
|
|
|
+ routes=set(state["route_decision"].routes)
|
|
|
+ if RouteName.CLARIFY in routes:
|
|
|
+ return "clarify"
|
|
|
+ if RouteName.REFUSE in routes:
|
|
|
+ return "refuse"
|
|
|
+ if RouteName.DIRECT_ANSWER in routes:
|
|
|
+ return "generate"
|
|
|
+ return "plan"
|
|
|
+
|
|
|
+ def plan_query(state:AgentState)->dict:
|
|
|
+ plan=builder.engine.plan(state["current_query"],state["route_decision"])
|
|
|
+ return{
|
|
|
+ "retrieval_plan":plan,
|
|
|
+ "trace":trace_event(
|
|
|
+ state,
|
|
|
+ node="plan_query",
|
|
|
+ detail={
|
|
|
+ "steps":plan.steps
|
|
|
+ }
|
|
|
+ )
|
|
|
+ }
|
|
|
+
|
|
|
+ def execute_step(step:PlanStep)->ToolResult:
|
|
|
+ if step.tool==RouteName.MILVUS_SEARCH:
|
|
|
+ return builder.vector_tool.invoke(
|
|
|
+ query=step.query,
|
|
|
+ service=step.arguments.get("service",""),
|
|
|
+ fault_type=step.arguments.get("fault_type",""),
|
|
|
+ client_os=step.arguments.get("client_os",""),
|
|
|
+ top_k=step.arguments.get("top_k",5)
|
|
|
+ )
|
|
|
+ if step.tool==RouteName.SQL_QUERY:
|
|
|
+ return builder.sql_tool.invoke(
|
|
|
+ service=step.arguments.get("service",""),
|
|
|
+ fault_type=step.arguments.get("fault_type",""),
|
|
|
+ days=step.arguments.get("days",30),
|
|
|
+ client_os=step.arguments.get("client_os","")
|
|
|
+ )
|
|
|
+ if step.tool==RouteName.WEB_SEARCH:
|
|
|
+ return builder.web_tool.invoke(
|
|
|
+ query=step.query,
|
|
|
+ max_results=step.arguments.get("max_results",5)
|
|
|
+ )
|
|
|
+ return ToolResult(
|
|
|
+ status="error",
|
|
|
+ tool=step.tool.value,
|
|
|
+ error_code="UNSUPPORTED_TOOL",
|
|
|
+ error_message=f"未注册工具:{step.tool.value}",
|
|
|
+ )
|
|
|
+
|
|
|
+ def execute_plan(state:AgentState)->dict:
|
|
|
+ steps=state["retrieval_plan"].steps
|
|
|
+ if not steps:
|
|
|
+ results:list[ToolResult]=[]
|
|
|
+ else:
|
|
|
+ with ThreadPoolExecutor(max_workers=min(len(steps),4)) as executor:
|
|
|
+ results=list(executor.map(execute_step, steps))
|
|
|
+ evidence=[]
|
|
|
+ for result in results:
|
|
|
+ evidence.extend(result.evidence)
|
|
|
+ errors = [
|
|
|
+ {
|
|
|
+ "tool": result.tool,
|
|
|
+ "code": result.error_code,
|
|
|
+ "message": result.error_message,
|
|
|
+ }
|
|
|
+ for result in results
|
|
|
+ if result.status == "error"
|
|
|
+ ]
|
|
|
+ executed_queries = [
|
|
|
+ {
|
|
|
+ "step_id": step.id,
|
|
|
+ "tool": result.tool,
|
|
|
+ "query": step.query,
|
|
|
+ "arguments": step.arguments,
|
|
|
+ "status": result.status,
|
|
|
+ "latency_ms": result.latency_ms,
|
|
|
+ }
|
|
|
+ for step, result in zip(steps, results, strict=True)
|
|
|
+ ]
|
|
|
+ return {
|
|
|
+ "tool_results": results,
|
|
|
+ "executed_queries": executed_queries,
|
|
|
+ "evidence": evidence,
|
|
|
+ "errors": [*state.get("errors", []), *errors],
|
|
|
+ "trace": trace_event(
|
|
|
+ state,
|
|
|
+ "execute_plan",
|
|
|
+ {
|
|
|
+ "tools": [result.tool for result in results],
|
|
|
+ "statuses": [result.status for result in results],
|
|
|
+ "evidence_count": len(evidence),
|
|
|
+ "executed_queries": executed_queries,
|
|
|
+ },
|
|
|
+ ),
|
|
|
+ }
|
|
|
+
|
|
|
+ def grade_evidence(state:AgentState)->dict:
|
|
|
+ current_round = state.get("retrieval_round", 0)
|
|
|
+ max_rounds = state["route_decision"].max_rounds
|
|
|
+ grade=builder.engine.grade(
|
|
|
+ query=state["current_query"],
|
|
|
+ decision=state["route_decision"],
|
|
|
+ evidence=state["evidence"],
|
|
|
+ current_round=current_round,
|
|
|
+ max_rounds=max_rounds,
|
|
|
+ min_score=builder.settings.min_evidence_score
|
|
|
+ )
|
|
|
+ if(grade.recommended_action=="rewrite_query"
|
|
|
+ and current_round >= max_rounds):
|
|
|
+ grade = grade.model_copy(
|
|
|
+ update={
|
|
|
+ "sufficient": False,
|
|
|
+ "recommended_action": "stop",
|
|
|
+ "reason": (
|
|
|
+ f"{grade.reason};已达到最大检索轮数 {max_rounds},"
|
|
|
+ "程序层强制停止。"
|
|
|
+ ),
|
|
|
+ }
|
|
|
+ )
|
|
|
+ return {
|
|
|
+ "quality_grade":grade,
|
|
|
+ "trace":trace_event(
|
|
|
+ state=state,
|
|
|
+ node="grade_evidence",
|
|
|
+ detail=grade.model_dump(mode="json"),
|
|
|
+ ),
|
|
|
+ }
|
|
|
+
|
|
|
+ def after_grade(state:AgentState)->str:
|
|
|
+ action=state["quality_grade"].recommended_action
|
|
|
+ if action=="accept":
|
|
|
+ return "generate"
|
|
|
+ if action=="rewrite_query":
|
|
|
+ return "rewrite"
|
|
|
+ return "generate_partial"
|
|
|
+
|
|
|
+ def rewrite_query(state:AgentState)->dict:
|
|
|
+ rewritten=builder.engine.rewrite(
|
|
|
+ query=state["current_query"],
|
|
|
+ grade=state["quality_grade"])
|
|
|
+ next_round=state.get("retrieval_round",0)+1
|
|
|
+ return{
|
|
|
+ "current_query":rewritten,
|
|
|
+ "retrieval_round":next_round,
|
|
|
+ "trace":trace_event(state=state,
|
|
|
+ node="rewrite_query",
|
|
|
+ detail={
|
|
|
+ "round":next_round,
|
|
|
+ "rewritten":rewritten
|
|
|
+ })
|
|
|
+ }
|
|
|
+
|
|
|
+ def generate_answer(state: AgentState) -> dict:
|
|
|
+ answer = builder.engine.answer(
|
|
|
+ state["original_query"], state.get("evidence", []), partial=False
|
|
|
+ )
|
|
|
+ return {
|
|
|
+ "final_answer": answer,
|
|
|
+ "termination_reason": "evidence_accepted"
|
|
|
+ if state.get("evidence")
|
|
|
+ else "direct_answer",
|
|
|
+ "trace": trace_event(state, "generate_answer", {"partial": False}),
|
|
|
+ }
|
|
|
+
|
|
|
+ def generate_partial_answer(state: AgentState) -> dict:
|
|
|
+ answer = builder.engine.answer(
|
|
|
+ state["original_query"], state.get("evidence", []), partial=True
|
|
|
+ )
|
|
|
+ return {
|
|
|
+ "final_answer": answer,
|
|
|
+ "termination_reason": "retrieval_budget_exhausted",
|
|
|
+ "trace": trace_event(state, "generate_partial_answer", {"partial": True}),
|
|
|
+ }
|
|
|
+
|
|
|
+ def clarify(state: AgentState) -> dict:
|
|
|
+ return {
|
|
|
+ "final_answer": "当前流程仅处理阿里云 SSL-VPN 的认证超时故障;请确认是否出现认证超时提示,并补充客户端系统或版本信息。",
|
|
|
+ "termination_reason": "clarification_required",
|
|
|
+ "trace": trace_event(state, "clarify", {}),
|
|
|
+ }
|
|
|
+
|
|
|
+ def refuse(state: AgentState) -> dict:
|
|
|
+ return {
|
|
|
+ "final_answer": "当前请求涉及受限数据,系统拒绝执行。",
|
|
|
+ "termination_reason": "security_policy",
|
|
|
+ "trace": trace_event(state, "refuse", {}),
|
|
|
+ }
|
|
|
+
|
|
|
+ graph=StateGraph(AgentState)
|
|
|
+ graph.add_node("normalize_query",normalize_query)
|
|
|
+ graph.add_node("route_query",route_query)
|
|
|
+ graph.add_node("plan_query",plan_query)
|
|
|
+ graph.add_node("execute_plan",execute_plan)
|
|
|
+ graph.add_node("grade_evidence",grade_evidence)
|
|
|
+ graph.add_node("rewrite_query",rewrite_query)
|
|
|
+ graph.add_node("generate_answer",generate_answer)
|
|
|
+ graph.add_node("generate_partial_answer",generate_partial_answer)
|
|
|
+ graph.add_node("clarify",clarify)
|
|
|
+ graph.add_node("refuse",refuse)
|
|
|
+
|
|
|
+ graph.add_edge(START,"normalize_query")
|
|
|
+ graph.add_edge("normalize_query","route_query")
|
|
|
+ graph.add_conditional_edges(
|
|
|
+ "route_query",
|
|
|
+ after_route,
|
|
|
+ {
|
|
|
+ "clarify": "clarify",
|
|
|
+ "refuse": "refuse",
|
|
|
+ "generate": "generate_answer",
|
|
|
+ "plan": "plan_query",
|
|
|
+ }
|
|
|
+ )
|
|
|
+ graph.add_edge("plan_query","execute_plan")
|
|
|
+ graph.add_edge("execute_plan","grade_evidence")
|
|
|
+ graph.add_conditional_edges(
|
|
|
+ "grade_evidence",
|
|
|
+ after_grade,
|
|
|
+ {
|
|
|
+ "generate":"generate_answer",
|
|
|
+ "rewrite":"rewrite_query",
|
|
|
+ "generate_partial":"generate_partial_answer"
|
|
|
+ }
|
|
|
+ )
|
|
|
+
|
|
|
+ graph.add_edge("rewrite_query","plan_query")
|
|
|
+ graph.add_edge("generate_answer",END)
|
|
|
+ graph.add_edge("generate_partial_answer",END)
|
|
|
+ graph.add_edge("clarify", END)
|
|
|
+ graph.add_edge("refuse", END)
|
|
|
+
|
|
|
+ return graph.compile()
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|
|
|
+
|