Files
more_dots/agent/core/base_agent.py
T
2026-03-24 18:07:22 +08:00

89 lines
3.7 KiB
Python

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,
}