init
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
"""代理实现模块 - 具体的代理实现"""
|
||||
|
||||
from .conversation import ConversationAgent
|
||||
from .tool import ToolAgent
|
||||
|
||||
__all__ = ["ConversationAgent", "ToolAgent"]
|
||||
@@ -0,0 +1,112 @@
|
||||
from typing import Dict, Any, Optional, cast
|
||||
from langchain_core.messages import HumanMessage, AIMessage
|
||||
from langgraph.graph import StateGraph, END
|
||||
|
||||
from agent.core.base_agent import BaseAgent
|
||||
from agent.core.state import AgentState
|
||||
from agent.core import nodes
|
||||
from config import CONVERSATION_MAX_HISTORY_MESSAGES
|
||||
|
||||
|
||||
class ConversationAgent(BaseAgent):
|
||||
"""处理多轮对话的代理"""
|
||||
|
||||
def __init__(self, model_section: Optional[str] = None):
|
||||
super().__init__(model_section)
|
||||
|
||||
def _build_graph(self) -> Any:
|
||||
"""构建对话专用图"""
|
||||
workflow = StateGraph(cast(Any, AgentState))
|
||||
|
||||
workflow.add_node("analyze_intent", cast(Any, self._analyze_intent))
|
||||
self._add_shared_sql_nodes(workflow)
|
||||
workflow.add_node("update_context", cast(Any, self._update_context))
|
||||
|
||||
workflow.add_edge("analyze_intent", "process_input")
|
||||
self._add_shared_sql_edges(workflow, start_node="process_input", end_node="generate_response")
|
||||
workflow.add_edge("generate_response", "update_context")
|
||||
workflow.add_edge("update_context", END)
|
||||
|
||||
workflow.set_entry_point("analyze_intent")
|
||||
|
||||
return workflow.compile()
|
||||
|
||||
def _analyze_intent(self, state: AgentState) -> AgentState:
|
||||
"""分析用户意图与对话上下文"""
|
||||
user_message = state.messages[-1] if state.messages else None
|
||||
|
||||
if user_message and isinstance(user_message, HumanMessage):
|
||||
content = user_message.content.lower()
|
||||
|
||||
if any(word in content for word in ["hello", "hi", "hey", "greetings"]):
|
||||
state.intent = "greeting"
|
||||
elif any(word in content for word in ["help", "assist", "support"]):
|
||||
state.intent = "help"
|
||||
elif any(token in content for token in ["sql", "查询", "统计", "汇总", "top", "eta", "so ", "soid", "订单"]):
|
||||
state.intent = "sql_query"
|
||||
elif "?" in content:
|
||||
state.intent = "question"
|
||||
else:
|
||||
state.intent = "general"
|
||||
|
||||
state.sync_context()
|
||||
state.set_current_step("intent_analyzed")
|
||||
return state
|
||||
|
||||
def _generate_response(self, state: AgentState) -> AgentState:
|
||||
"""优先返回 SQL 执行结果,其次返回生成 SQL,再回退到模型回复"""
|
||||
return nodes.generate_response(state, self.model)
|
||||
|
||||
def _update_context(self, state: AgentState) -> AgentState:
|
||||
"""更新当前会话的对话上下文与历史。"""
|
||||
conversation_history = list(state.context.get("conversation_history") or [])
|
||||
for message in state.messages:
|
||||
if isinstance(message, (HumanMessage, AIMessage)):
|
||||
conversation_history.append(message)
|
||||
|
||||
# 使用配置文件中的最大消息数限制
|
||||
max_messages = CONVERSATION_MAX_HISTORY_MESSAGES
|
||||
if len(conversation_history) > max_messages:
|
||||
conversation_history = conversation_history[-max_messages:]
|
||||
|
||||
state.context["conversation_history"] = conversation_history
|
||||
|
||||
state.sync_context()
|
||||
last_context = dict(state.context)
|
||||
last_context.pop("conversation_history", None)
|
||||
last_context.pop("last_context", None)
|
||||
state.context["last_context"] = last_context
|
||||
state.set_current_step("context_updated")
|
||||
return state
|
||||
|
||||
def run(self, user_input: str, **kwargs) -> Dict[str, Any]:
|
||||
"""运行对话,历史与上下文由调用方按会话维度传入。"""
|
||||
context = dict(kwargs)
|
||||
context["conversation_history"] = list(context.get("conversation_history") or [])
|
||||
context["last_context"] = dict(context.get("last_context") or {})
|
||||
initial_state = AgentState(
|
||||
messages=[HumanMessage(content=user_input)],
|
||||
context=context
|
||||
)
|
||||
|
||||
result = self.graph.invoke(initial_state)
|
||||
final_state = self._coerce_state(initial_state, result)
|
||||
|
||||
return {
|
||||
"messages": final_state.messages,
|
||||
"context": final_state.sync_context(),
|
||||
"conversation_history": list(final_state.context.get("conversation_history") or []),
|
||||
"final_step": final_state.current_step,
|
||||
}
|
||||
|
||||
def stream_run(self, user_input: str, **kwargs):
|
||||
"""流式运行对话;会话历史需由调用方显式传入。"""
|
||||
conversation_history = list(kwargs.get("conversation_history") or [])
|
||||
all_messages = conversation_history + [HumanMessage(content=user_input)]
|
||||
full_text = ""
|
||||
|
||||
for chunk in self.model.stream(all_messages):
|
||||
if hasattr(chunk, "content") and chunk.content:
|
||||
full_text += chunk.content
|
||||
yield chunk.content
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
from typing import Dict, Any, List, Optional, cast
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langchain_core.tools import BaseTool
|
||||
from langgraph.graph import StateGraph, END
|
||||
from langgraph.prebuilt import ToolNode
|
||||
|
||||
from agent.core.base_agent import BaseAgent
|
||||
from agent.core.state import AgentState
|
||||
from tools.calculator import CalculatorTool
|
||||
from tools.web_search import WebSearchTool
|
||||
from tools.rest_api_tool import RestApiTool
|
||||
from tools.sr_api_tool import SrApiQueryTool
|
||||
|
||||
|
||||
class ToolAgent(BaseAgent):
|
||||
"""可使用工具完成任务的代理"""
|
||||
|
||||
def __init__(self, model_section: Optional[str] = None, tools: List[BaseTool] = None):
|
||||
if tools is None:
|
||||
tools = [CalculatorTool(), WebSearchTool(), RestApiTool(), SrApiQueryTool()]
|
||||
|
||||
self.tools = tools
|
||||
self.tool_node = ToolNode(tools)
|
||||
super().__init__(model_section)
|
||||
|
||||
def _build_graph(self) -> Any:
|
||||
"""构建可使用工具的图"""
|
||||
workflow = StateGraph(cast(Any, AgentState))
|
||||
|
||||
self._add_shared_sql_nodes(workflow)
|
||||
workflow.add_node("agent", cast(Any, self._agent_node))
|
||||
workflow.add_node("tools", cast(Any, self.tool_node))
|
||||
|
||||
workflow.set_entry_point("process_input")
|
||||
self._add_shared_sql_edges(workflow, start_node="process_input", end_node="agent")
|
||||
workflow.add_edge("tools", "agent")
|
||||
|
||||
workflow.add_conditional_edges(
|
||||
"agent",
|
||||
self._should_use_tools,
|
||||
{
|
||||
"tools": "tools",
|
||||
"end": END,
|
||||
}
|
||||
)
|
||||
|
||||
return cast(Any, workflow.compile())
|
||||
|
||||
def _agent_node(self, state: AgentState) -> AgentState:
|
||||
"""决定是否调用工具的代理节点"""
|
||||
model_with_tools = self.model.bind_tools(self.tools)
|
||||
|
||||
if state.messages:
|
||||
try:
|
||||
response = model_with_tools.invoke(state.messages)
|
||||
except Exception as e:
|
||||
error_text = str(e)
|
||||
if "tool choice" in error_text and "auto" in error_text:
|
||||
fallback_model = self.model.bind_tools(self.tools, tool_choice="none")
|
||||
response = fallback_model.invoke(state.messages)
|
||||
else:
|
||||
raise
|
||||
state.messages.append(response)
|
||||
|
||||
return state
|
||||
|
||||
def _should_use_tools(self, state: AgentState) -> str:
|
||||
"""判断是否需要使用工具"""
|
||||
last_message = state.messages[-1]
|
||||
|
||||
if hasattr(last_message, 'tool_calls') and last_message.tool_calls:
|
||||
return "tools"
|
||||
|
||||
return "end"
|
||||
|
||||
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(),
|
||||
"tools_used": [tool.name for tool in self.tools],
|
||||
"final_step": final_state.current_step,
|
||||
}
|
||||
Reference in New Issue
Block a user