| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291 |
- 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()
-
-
|