from typing import Any, Dict, Optional, cast from langchain_core.messages import HumanMessage from langgraph.graph import StateGraph, END from services.core.llm_factory import create_chat_model from .state import AgentState from . import nodes class BaseAgent: """包含通用功能的基础代理类""" def __init__(self, model_section: Optional[str] = None): self.model = create_chat_model(model_section) self.graph = self._build_graph() def _build_graph(self) -> Any: """构建代理状态图""" workflow = StateGraph(cast(Any, AgentState)) self._add_shared_sql_nodes(workflow) self._add_shared_sql_edges(workflow, start_node="process_input", end_node="generate_response") workflow.add_edge("generate_response", END) workflow.set_entry_point("process_input") return workflow.compile() def _add_shared_sql_nodes(self, workflow: Any) -> None: """注册 SQL 规划相关共享节点。""" workflow.add_node("process_input", cast(Any, nodes.process_input)) workflow.add_node("normalize_input", cast(Any, self._normalize_input)) workflow.add_node("classify_query_mode", cast(Any, nodes.classify_query_mode)) workflow.add_node("match_table", cast(Any, nodes.match_table)) workflow.add_node("load_sql_prompt", cast(Any, nodes.load_sql_prompt)) workflow.add_node("build_sql_plan", cast(Any, nodes.build_sql_plan)) workflow.add_node("generate_sql", cast(Any, self._generate_sql)) workflow.add_node("execute_sql", cast(Any, nodes.execute_sql)) workflow.add_node("check_empty_result", cast(Any, nodes.check_empty_result)) workflow.add_node("generate_response", cast(Any, self._generate_response)) @staticmethod def _add_shared_sql_edges(workflow: Any, start_node: str, end_node: str) -> None: """串联标准 SQL 工作流。""" workflow.add_edge(start_node, "normalize_input") workflow.add_edge("normalize_input", "classify_query_mode") workflow.add_edge("classify_query_mode", "match_table") workflow.add_edge("match_table", "load_sql_prompt") workflow.add_edge("load_sql_prompt", "build_sql_plan") workflow.add_edge("build_sql_plan", "generate_sql") workflow.add_edge("generate_sql", "execute_sql") workflow.add_edge("execute_sql", "check_empty_result") workflow.add_edge("check_empty_result", end_node) def _generate_response(self, state: AgentState) -> AgentState: """使用 LLM 生成回复""" return nodes.generate_response(state, self.model) def _normalize_input(self, state: AgentState) -> AgentState: """规范化用户输入""" return nodes.normalize_input(state, self.model) def _generate_sql(self, state: AgentState) -> AgentState: """生成 SQL""" return nodes.generate_sql(state, self.model) @staticmethod def _coerce_state(initial_state: AgentState, result: Any) -> AgentState: """兼容 LangGraph 返回 AgentState 或 dict。""" if isinstance(result, AgentState): return result return initial_state.apply_graph_result(result) def run(self, user_input: str, **kwargs) -> Dict[str, Any]: """运行代理并处理用户输入""" initial_state = AgentState( messages=[HumanMessage(content=user_input)], context=kwargs ) result = self.graph.invoke(initial_state) final_state = self._coerce_state(initial_state, result) return { "messages": final_state.messages, "context": final_state.sync_context(), "final_step": final_state.current_step, }