This commit is contained in:
2026-02-26 13:43:44 +08:00
parent 2c2db92ae9
commit 68200cdfe6
51 changed files with 2107 additions and 351 deletions
+6
View File
@@ -0,0 +1,6 @@
from .state import AgentState
from .graph import BaseAgent
from .conversation import ConversationAgent
from .tool import ToolAgent
__all__ = ["AgentState", "BaseAgent", "ConversationAgent", "ToolAgent"]
+106
View File
@@ -0,0 +1,106 @@
from typing import Dict, Any, List, Optional
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage
from langgraph.graph import StateGraph, END
from .graph import BaseAgent
from .state import AgentState
class ConversationAgent(BaseAgent):
"""处理多轮对话的代理"""
def __init__(self, model_section: Optional[str] = None):
super().__init__(model_section)
self.conversation_history: List[BaseMessage] = []
def _build_graph(self) -> StateGraph:
"""构建对话专用图"""
workflow = StateGraph(AgentState)
workflow.add_node("analyze_intent", self._analyze_intent)
workflow.add_node("normalize_input", self._normalize_input)
workflow.add_node("generate_response", self._generate_response)
workflow.add_node("update_context", self._update_context)
workflow.add_edge("analyze_intent", "normalize_input")
workflow.add_edge("normalize_input", "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.context["intent"] = "greeting"
elif any(word in content for word in ["help", "assist", "support"]):
state.context["intent"] = "help"
elif "?" in content:
state.context["intent"] = "question"
else:
state.context["intent"] = "general"
state.current_step = "intent_analyzed"
return state
def _generate_response(self, state: AgentState) -> AgentState:
"""结合对话历史生成回复"""
all_messages = self.conversation_history + state.messages
if all_messages:
response = self.model.invoke(all_messages)
state.messages.append(response)
state.current_step = "response_generated"
return state
def _update_context(self, state: AgentState) -> AgentState:
"""更新对话上下文与历史"""
for message in state.messages:
if isinstance(message, (HumanMessage, AIMessage)):
self.conversation_history.append(message)
if len(self.conversation_history) > 10:
self.conversation_history = self.conversation_history[-10:]
state.current_step = "context_updated"
return state
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)
return {
"messages": result.get("messages", []),
"context": result.get("context", {}),
"conversation_history": self.conversation_history,
"final_step": result.get("current_step", "unknown")
}
def stream_run(self, user_input: str):
"""流式运行对话并维护历史"""
all_messages = self.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
self.conversation_history.append(HumanMessage(content=user_input))
self.conversation_history.append(AIMessage(content=full_text))
if len(self.conversation_history) > 10:
self.conversation_history = self.conversation_history[-10:]
+54
View File
@@ -0,0 +1,54 @@
from typing import Any, Dict, Optional
from langchain_core.messages import HumanMessage
from langgraph.graph import StateGraph, END
from services.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) -> StateGraph:
"""构建代理状态图"""
workflow = StateGraph(AgentState)
workflow.add_node("process_input", nodes.process_input)
workflow.add_node("normalize_input", self._normalize_input)
workflow.add_node("generate_response", self._generate_response)
workflow.add_edge("process_input", "normalize_input")
workflow.add_edge("normalize_input", "generate_response")
workflow.add_edge("generate_response", END)
workflow.set_entry_point("process_input")
return workflow.compile()
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 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)
return {
"messages": result.get("messages", []),
"context": result.get("context", {}),
"final_step": result.get("current_step", "unknown")
}
+43
View File
@@ -0,0 +1,43 @@
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage
from .state import AgentState
from services.prompt_manager import PromptManager
from services.template_matcher import TemplateMatcher
def process_input(state: AgentState) -> AgentState:
"""处理用户输入"""
state.current_step = "processed"
return state
def generate_response(state: AgentState, model) -> AgentState:
"""使用 LLM 生成回复"""
if state.messages:
response = model.invoke(state.messages)
state.messages.append(response)
return state
def normalize_input(state: AgentState, model) -> AgentState:
"""将用户输入规范化为标准英文语句"""
if not state.messages:
return state
last_message = state.messages[-1]
if not isinstance(last_message, HumanMessage):
return state
prompt_manager = PromptManager()
system_prompt = SystemMessage(
content=prompt_manager.get("system", "english_normalizer")
)
response = model.invoke([system_prompt, HumanMessage(content=last_message.content)])
normalized = response.content if hasattr(response, "content") else str(response)
state.context["original_input"] = last_message.content
state.context["normalized_input"] = normalized
matcher = TemplateMatcher()
state.context["table_match"] = matcher.match(normalized)
return state
+14
View File
@@ -0,0 +1,14 @@
from typing import Any, Dict, List
from langchain_core.messages import BaseMessage
class AgentState:
"""代理工作流的状态定义"""
messages: List[BaseMessage]
current_step: str
context: Dict[str, Any]
def __init__(self, messages: List[BaseMessage] = None, current_step: str = "start", context: Dict[str, Any] = None):
self.messages = messages or []
self.current_step = current_step
self.context = context or {}
+91
View File
@@ -0,0 +1,91 @@
from typing import Dict, Any, List, Optional
from langchain_core.messages import BaseMessage, HumanMessage
from langchain_core.tools import BaseTool
from langgraph.graph import StateGraph, END
from langgraph.prebuilt import ToolNode
from .graph import BaseAgent
from .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) -> StateGraph:
"""构建可使用工具的图"""
workflow = StateGraph(AgentState)
workflow.add_node("normalize_input", self._normalize_input)
workflow.add_node("agent", self._agent_node)
workflow.add_node("tools", self.tool_node)
workflow.add_edge("normalize_input", "agent")
workflow.add_edge("tools", "agent")
workflow.add_conditional_edges(
"agent",
self._should_use_tools,
{
"tools": "tools",
"end": END,
}
)
workflow.set_entry_point("normalize_input")
return 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)
return {
"messages": result.get("messages", []),
"context": result.get("context", {}),
"tools_used": [tool.name for tool in self.tools],
"final_step": result.get("current_step", "unknown")
}
+1
View File
@@ -0,0 +1 @@
# Agent 内部工具函数(按需扩展)