From f26ca077167cac49358dc8df7a13388c7d838450 Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Thu, 4 Sep 2025 16:05:08 -0400 Subject: [PATCH] minimalistic tool registration --- libs/langgraph/langgraph/agent/__init__.py | 15 +++++++++++---- .../langgraph/langgraph/agent/middleware/swarm.py | 5 +++-- .../langgraph/agent/middleware/tool_calls.py | 11 ++++++++++- libs/langgraph/langgraph/agent/types.py | 3 +++ 4 files changed, 27 insertions(+), 7 deletions(-) diff --git a/libs/langgraph/langgraph/agent/__init__.py b/libs/langgraph/langgraph/agent/__init__.py index 3b428d759..9afec3a6b 100644 --- a/libs/langgraph/langgraph/agent/__init__.py +++ b/libs/langgraph/langgraph/agent/__init__.py @@ -43,8 +43,14 @@ def create_agent( model = cast(BaseChatModel, init_chat_model(model)) - # init tool node - tool_node = tools if isinstance(tools, ToolNode) else ToolNode(tools=tools) + if isinstance(tools, list): + all_tools = [t for m in middleware for t in m.tools] + list(tools) + tool_node = ToolNode(tools=all_tools) + default_tools = tools + else: + # TODO: what do we do when middleware tools are specified, plus a ToolNode is used? + tool_node = tools + default_tools = list(tool_node.tools_by_name.values()) # validate middleware assert len({m.__class__.__name__ for m in middleware}) == len(middleware), ( @@ -77,7 +83,7 @@ def create_agent( def model_request(state: AgentState) -> AgentState: request = state.model_request or ModelRequest( model=model, - tools=list(tool_node.tools_by_name.values()), + tools=default_tools, system_prompt=system_prompt, response_format=response_format, messages=state.messages, @@ -85,6 +91,7 @@ def create_agent( ) # prepare messages + print(request.system_prompt) if request.system_prompt: messages = [SystemMessage(request.system_prompt)] + request.messages else: @@ -127,7 +134,7 @@ def create_agent( default_model_request = ModelRequest( model=model, - tools=list(tool_node.tools_by_name.values()), + tools=default_tools, system_prompt=system_prompt, response_format=response_format, messages=state.messages, diff --git a/libs/langgraph/langgraph/agent/middleware/swarm.py b/libs/langgraph/langgraph/agent/middleware/swarm.py index ad55280b6..364bf8e8f 100644 --- a/libs/langgraph/langgraph/agent/middleware/swarm.py +++ b/libs/langgraph/langgraph/agent/middleware/swarm.py @@ -15,8 +15,8 @@ class SwarmAgent: tools: list[BaseTool] -class SwarmMiddleware(AgentMiddleware): - """Swarm middleware. +class MultiAgentMiddleware(AgentMiddleware): + """Multi agent middleware (enabling swarm like behavior). TODOs: * Support create_agent for handoffs @@ -59,6 +59,7 @@ class SwarmMiddleware(AgentMiddleware): request.tools = agent.tools request.tools.extend(self.handoff_tools) + return request def after_model(self, state) -> State | None: diff --git a/libs/langgraph/langgraph/agent/middleware/tool_calls.py b/libs/langgraph/langgraph/agent/middleware/tool_calls.py index e55d4b9ee..7e4cee6ad 100644 --- a/libs/langgraph/langgraph/agent/middleware/tool_calls.py +++ b/libs/langgraph/langgraph/agent/middleware/tool_calls.py @@ -2,6 +2,7 @@ from dataclasses import dataclass, field from typing import Annotated, Any, Dict, List, cast from langchain_core.messages import AIMessage +from typing_extensions import Annotated from langgraph.agent.types import AgentJump, AgentMiddleware, AgentState, AgentUpdate @@ -11,7 +12,15 @@ class ToolCallLimitMiddleware(AgentMiddleware): @dataclass class State(AgentMiddleware.State): - tool_call_count: dict[str, int] = field(default_factory=dict) + important: Annotated[dict[str, int], Input, Output] = field(default_factory=dict) + + @dataclass + class InputState(AgentMiddleware.State): + important: dict[str, int] + + @dataclass + class OutputState(AgentMiddleware.State): + important: dict[str, int] def __init__(self, tool_limits: dict[str, int]): self.tool_limits = tool_limits diff --git a/libs/langgraph/langgraph/agent/types.py b/libs/langgraph/langgraph/agent/types.py index 2dc139692..32c457e39 100644 --- a/libs/langgraph/langgraph/agent/types.py +++ b/libs/langgraph/langgraph/agent/types.py @@ -12,6 +12,7 @@ from typing_extensions import TypedDict from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.graph.message import Messages, add_messages +from langgraph.runtime import Runtime ResponseFormat = dict | type[BaseModel] JumpTo = Literal["tools", "model", "__end__"] @@ -39,6 +40,8 @@ class AgentMiddleware: class State(AgentState): pass + tools: list[BaseTool] + def before_model(self, state: State) -> AgentUpdate | AgentJump | None: pass