from __future__ import annotations from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from datetime import datetime, timezone from typing import Literal from langgraph.graph import END, START, StateGraph from app.config import Settings from app.decision_engine import DecisionEngine from app.schemas import RouteName, ToolResult from app.state import AgentState from app.tools import SQLQueryTool, VectorSearchTool, WebSearchTool def trace_event(state: AgentState, node: str, detail: dict) -> list[dict]: return [ *state.get("trace", []), { "node": node, "at": datetime.now(timezone.utc).isoformat(), "detail": detail, }, ] @dataclass class GraphDependencies: settings: Settings engine: DecisionEngine vector_tool: VectorSearchTool sql_tool: SQLQueryTool web_tool: WebSearchTool def build_graph(deps: GraphDependencies): 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 = deps.engine.route( state["current_query"], deps.settings.max_retrieval_rounds ) return { "route_decision": decision, "trace": trace_event( state, "route_query", { "intent": decision.intent, "routes": [route.value for route in decision.routes], "reason_code": decision.reason_code, }, ), } def after_route( state: AgentState, ) -> Literal["clarify", "refuse", "generate", "plan"]: 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 = deps.engine.plan(state["current_query"], state["route_decision"]) return { "retrieval_plan": plan, "trace": trace_event( state, "plan_query", {"steps": [step.model_dump(mode="json") for step in plan.steps]}, ), } def execute_step(step) -> ToolResult: # Executor 只识别注册过的枚举工具;模型不能构造任意函数名。 if step.tool == RouteName.MILVUS_SEARCH: return deps.vector_tool.invoke( query=step.query, policy_type=str(step.arguments.get("policy_type", "")), top_k=int(step.arguments.get("top_k", 5)), ) if step.tool == RouteName.SQL_QUERY: return deps.sql_tool.invoke( user_id=str(step.arguments.get("user_id", "")), days=int(step.arguments.get("days", 30)), product_keyword=str(step.arguments.get("product_keyword", "")), ) if step.tool == RouteName.WEB_SEARCH: requested = int(step.arguments.get("max_results", 5)) # 即使模型给出更大的 max_results,也不能突破程序配置上限。 return deps.web_tool.invoke( query=step.query, max_results=min(max(requested, 1), deps.settings.web_search_max_results), ) 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] = [] elif len(steps) == 1: results = [execute_step(steps[0])] else: # 本项目的计划均为无依赖读任务;依赖型 DAG 需按拓扑批次调度。 with ThreadPoolExecutor(max_workers=min(len(steps), 4)) as executor: results = list(executor.map(execute_step, steps)) # ToolResult 用于控制流,Evidence 用于答案生成,二者不能混为一体。 evidence = [item for result in results for item in result.evidence] errors = [ { "tool": result.tool, "code": result.error_code, "message": result.error_message, } for result in results if result.status == "error" ] # 将规划 Query、受控参数和实际结果并列记录,便于观察真实调用。 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 = deps.engine.grade( query=state["current_query"], decision=state["route_decision"], evidence=state.get("evidence", []), current_round=current_round, max_rounds=max_rounds, min_score=deps.settings.min_evidence_score, ) # 模型只能建议 Rewrite,是否还有预算必须由程序层决定。 # 达到上限后强制 Stop,避免模型持续 Rewrite 导致图无限循环。 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, "grade_evidence", grade.model_dump(mode="json"), ), } def after_grade(state: AgentState) -> Literal["generate", "rewrite", "generate_partial"]: # 条件边只消费结构化动作,不解析 grader 的自然语言 reason。 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 = deps.engine.rewrite( state["current_query"], state["quality_grade"] ) next_round = state.get("retrieval_round", 0) + 1 return { "current_query": rewritten, "retrieval_round": next_round, "trace": trace_event( state, "rewrite_query", {"round": next_round, "rewritten_query": rewritten}, ), } def generate_answer(state: AgentState) -> dict: answer = deps.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 = deps.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": "请补充用户编号、商品名称或订单范围后再查询。", "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", {}), } builder = StateGraph(AgentState) builder.add_node("normalize_query", normalize_query) builder.add_node("route_query", route_query) builder.add_node("plan_query", plan_query) builder.add_node("execute_plan", execute_plan) builder.add_node("grade_evidence", grade_evidence) builder.add_node("rewrite_query", rewrite_query) builder.add_node("generate_answer", generate_answer) builder.add_node("generate_partial_answer", generate_partial_answer) builder.add_node("clarify", clarify) builder.add_node("refuse", refuse) builder.add_edge(START, "normalize_query") builder.add_edge("normalize_query", "route_query") builder.add_conditional_edges( "route_query", after_route, { "clarify": "clarify", "refuse": "refuse", "generate": "generate_answer", "plan": "plan_query", }, ) builder.add_edge("plan_query", "execute_plan") builder.add_edge("execute_plan", "grade_evidence") builder.add_conditional_edges( "grade_evidence", after_grade, { "generate": "generate_answer", "rewrite": "rewrite_query", "generate_partial": "generate_partial_answer", }, ) builder.add_edge("rewrite_query", "plan_query") builder.add_edge("generate_answer", END) builder.add_edge("generate_partial_answer", END) builder.add_edge("clarify", END) builder.add_edge("refuse", END) return builder.compile()