init
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
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,
|
||||
}
|
||||
Reference in New Issue
Block a user