minimalistic tool registration

This commit is contained in:
Sydney Runkle
2025-09-04 16:05:08 -04:00
parent b250823532
commit f26ca07716
4 changed files with 27 additions and 7 deletions
+11 -4
View File
@@ -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,
@@ -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:
@@ -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
+3
View File
@@ -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