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