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