Compare commits

...
Author SHA1 Message Date
Sydney Runkle f26ca07716 minimalistic tool registration 2025-09-04 16:05:08 -04:00
Sydney Runkle b250823532 tool and model calls 2025-09-04 10:15:03 -04:00
Sydney Runkle 120d34303d adding model calls 2025-09-04 09:59:45 -04:00
Sydney Runkle 6ff9e4a764 limiting calls 2025-09-04 09:43:53 -04:00
Sydney Runkle 3c36d2e2c8 initial test for swarm 2025-09-03 15:27:13 -04:00
Sydney Runkle f6d0382d66 more swarm progress 2025-09-03 15:14:33 -04:00
Sydney Runkle 0386fe5f6a swarm 2025-09-03 13:43:47 -04:00
Sydney Runkle b11ece823b first pass at modify as new node 2025-09-03 13:06:26 -04:00
9 changed files with 337 additions and 107 deletions
+106 -59
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), (
@@ -55,6 +61,11 @@ def create_agent(
for m in middleware
if m.__class__.before_model is not AgentMiddleware.before_model
]
middleware_w_modify_model_request = [
m
for m in middleware
if m.__class__.modify_model_request is not AgentMiddleware.modify_model_request
]
middleware_w_after = [
m
for m in middleware
@@ -68,17 +79,47 @@ def create_agent(
output_schema=AgentUpdate,
context_schema=context_schema,
)
graph.add_node(
"model_request",
_make_model_request_node(
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,
middleware=middleware,
response_format=response_format,
),
)
messages=state.messages,
tool_choice=None,
)
# prepare messages
print(request.system_prompt)
if request.system_prompt:
messages = [SystemMessage(request.system_prompt)] + request.messages
else:
messages = request.messages
# call model
if request.response_format:
model_ = request.model.with_structured_output(
request.response_format, include_raw=True
)
output = model_.invoke(
messages, tools=request.tools, tool_choice=request.tool_choice
)
return {"messages": output["raw"], "response": output["parsed"]}
else:
model_ = request.model.bind_tools(
request.tools,
tool_choice=request.tool_choice,
parallel_tool_calls=False,
)
output = model_.invoke(messages)
if state.response is not None:
return {"messages": output, "response": None}
else:
return {"messages": output}
graph.add_node("model_request", model_request)
graph.add_node("tools", tool_node)
for m in middleware:
if m.__class__.before_model is not AgentMiddleware.before_model:
graph.add_node(
@@ -86,6 +127,32 @@ def create_agent(
m.before_model,
input_schema=m.State,
)
if m.__class__.modify_model_request is not AgentMiddleware.modify_model_request:
def modify_model_request_node(state: AgentState) -> dict[str, ModelRequest]:
# TODO assert request.tools in tools, or pass them to tool node
default_model_request = ModelRequest(
model=model,
tools=default_tools,
system_prompt=system_prompt,
response_format=response_format,
messages=state.messages,
tool_choice=None,
)
return {
"model_request": m.modify_model_request(
state.model_request or default_model_request, state
)
}
graph.add_node(
f"{m.__class__.__name__}.modify_model_request",
modify_model_request_node,
input_schema=m.State,
)
if m.__class__.after_model is not AgentMiddleware.after_model:
graph.add_node(
f"{m.__class__.__name__}.after_model",
@@ -97,6 +164,8 @@ def create_agent(
first_node = (
f"{middleware_w_before[0].__class__.__name__}.before_model"
if middleware_w_before
else f"{middleware_w_modify_model_request[0].__class__.__name__}.modify_model_request"
if middleware_w_modify_model_request
else "model_request"
)
last_node = (
@@ -113,7 +182,7 @@ def create_agent(
[first_node, END],
)
graph.add_conditional_edges(
last_node, _make_model_to_tools_edge(first_node), ["tools", END]
last_node, _make_model_to_tools_edge(first_node), [first_node, "tools", END]
)
# add before model edges
@@ -134,6 +203,26 @@ def create_agent(
first_node,
)
# add modify model request edges
if middleware_w_modify_model_request:
for m1, m2 in zip(
middleware_w_modify_model_request, middleware_w_modify_model_request[1:]
):
_add_middleware_edge(
graph,
m1.modify_model_request,
f"{m1.__class__.__name__}.modify_model_request",
f"{m2.__class__.__name__}.modify_model_request",
first_node,
)
_add_middleware_edge(
graph,
middleware_w_modify_model_request[-1].modify_model_request,
f"{middleware_w_modify_model_request[-1].__class__.__name__}.modify_model_request",
"model_request",
first_node,
)
# add after model edges
if middleware_w_after:
graph.add_edge(
@@ -149,59 +238,17 @@ def create_agent(
f"{m2.__class__.__name__}.after_model",
first_node,
)
# _add_middleware_edge(
# graph,
# middleware_w_after[-1].after_model,
# f"{middleware_w_after[-1].__class__.__name__}.after_model",
# "model_request",
# first_node,
# )
return graph
def _make_model_request_node(
*,
system_prompt: str,
model: BaseChatModel,
tools: Sequence[BaseTool],
middleware: Sequence[AgentMiddleware] = (),
response_format: ResponseFormat | None = None,
) -> Callable[[AgentState], AgentState]:
def model_request(state: AgentState) -> AgentState:
# create request
request = ModelRequest(
model=model,
system_prompt=system_prompt,
messages=state.messages,
tool_choice=None,
tools=tools,
response_format=response_format,
)
# visit middleware in order
for mw in middleware:
request = mw.modify_model_request(request, state)
# TODO assert request.tools in tools, or pass them to tool node
# prepare messages
if request.system_prompt:
messages = [SystemMessage(request.system_prompt)] + request.messages
else:
messages = request.messages
# call model
if request.response_format:
model_ = request.model.with_structured_output(
request.response_format, include_raw=True
)
output = model_.invoke(
messages, tools=request.tools, tool_choice=request.tool_choice
)
return {"messages": output["raw"], "response": output["parsed"]}
else:
model_ = request.model
output = model_.invoke(
messages, tools=request.tools, tool_choice=request.tool_choice
)
if state.response is not None:
return {"messages": output, "response": None}
else:
return {"messages": output}
return model_request
def _resolve_jump(jump_to: JumpTo | None, first_node: str) -> str | None:
if jump_to == "model":
return first_node
@@ -0,0 +1,35 @@
from dataclasses import dataclass
from typing import Literal
from langchain_core.language_models.chat_models import BaseChatModel
from langgraph.agent.types import (
AgentJump,
AgentMiddleware,
AgentState,
AgentUpdate,
ModelRequest,
)
class DynamicModelMiddleware(AgentMiddleware):
"""Selects different models based on task complexity"""
def __init__(
self,
basic_model: BaseChatModel,
complex_model: BaseChatModel,
message_threshold: int = 5,
):
self.basic_model = basic_model
self.complex_model = complex_model
self.message_threshold = message_threshold
def modify_model_request(
self, request: ModelRequest, state: AgentState
) -> ModelRequest:
if len(state.messages) > self.message_threshold:
request.model = self.complex_model
else:
request.model = self.basic_model
return request
@@ -0,0 +1,23 @@
import operator
from dataclasses import dataclass
from typing import Annotated
from langgraph.agent.types import AgentJump, AgentMiddleware, AgentState, AgentUpdate
class ModelRequestLimitMiddleware(AgentMiddleware):
"""Terminates after N model requests"""
@dataclass
class State(AgentMiddleware.State):
model_request_count: Annotated[int, operator.add] = 0
def __init__(self, max_requests: int = 10):
self.max_requests = max_requests
def before_model(self, state: State) -> AgentUpdate | AgentJump | None:
# TODO: want to be able to configure end behavior here
if state.model_request_count == self.max_requests:
return {"jump_to": "__end__"}
return {"model_request_count": 1}
@@ -1,21 +1,22 @@
import uuid
from collections.abc import Sequence
from typing import Callable, Iterable
from langchain_core.language_models import LanguageModelLike
from langchain_core.messages import RemoveMessage, MessageLikeRepresentation
from langchain_core.messages import (
AIMessage,
AnyMessage,
MessageLikeRepresentation,
RemoveMessage,
ToolMessage,
)
from langchain_core.messages.utils import count_tokens_approximately
from collections.abc import Sequence
from langchain_core.messages import AnyMessage, AIMessage, ToolMessage
import uuid
TokenCounter = Callable[[Iterable[MessageLikeRepresentation]], int]
from langgraph.agent.types import AgentMiddleware, AgentState
DEFAULT_SUMMARY_PROMPT = """<role>
Context Extraction Assistant
</role>
@@ -41,13 +42,15 @@ Respond ONLY with the extracted context. Do not include any additional informati
class SummarizationMiddleware(AgentMiddleware):
def __init__(self,model: LanguageModelLike,
max_tokens_before_summary: int | None = None,
token_counter: TokenCounter = count_tokens_approximately,
messages_to_leave: int = 20,
summary_system_prompt: str = DEFAULT_SUMMARY_PROMPT,
fake_tool_call_name: str = "summarize_convo"):
def __init__(
self,
model: LanguageModelLike,
max_tokens_before_summary: int | None = None,
token_counter: TokenCounter = count_tokens_approximately,
messages_to_leave: int = 20,
summary_system_prompt: str = DEFAULT_SUMMARY_PROMPT,
fake_tool_call_name: str = "summarize_convo",
):
super().__init__()
self.model = model
self.max_tokens_before_summary = max_tokens_before_summary
@@ -64,33 +67,38 @@ class SummarizationMiddleware(AgentMiddleware):
return None
# Otherwise, we create a summary!
# Get messages that we want to create a summary for
messages_to_summarize = messages[:-self.messages_to_leave]
messages_to_summarize = messages[: -self.messages_to_leave]
# Create summary text
summary = self._summarize_messages(messages_to_summarize)
# Create fake messages to add to history
fake_tool_call_id = str(uuid.uuid4())
fake_messages = [AIMessage(
content="Looks like I'm running out of tokens. I'm going to summarize the conversation history to free up space.",
tool_calls={
"id": fake_tool_call_id,
"name": self.fake_tool_call_name,
"args": {
"reasoning":
"I'm running out of tokens. I'm going to summarize all of the messages since my last summary message to free up space.",
}
}),
ToolMessage(tool_call_id= fake_tool_call_id, content=summary)]
fake_messages = [
AIMessage(
content="Looks like I'm running out of tokens. I'm going to summarize the conversation history to free up space.",
tool_calls={
"id": fake_tool_call_id,
"name": self.fake_tool_call_name,
"args": {
"reasoning": "I'm running out of tokens. I'm going to summarize all of the messages since my last summary message to free up space.",
},
},
),
ToolMessage(tool_call_id=fake_tool_call_id, content=summary),
]
return {
"messages": [RemoveMessage(id=m.id) for m in messages_to_summarize] + fake_messages
"messages": [RemoveMessage(id=m.id) for m in messages_to_summarize]
+ fake_messages
}
def _summarize_messages(self, messages_to_summarize: Sequence[AnyMessage]) -> str:
system_message = self.summary_system_prompt
user_message = self._format_messages(messages_to_summarize)
response = self.model.invoke([
{"role": "system", "content": system_message},
{"role": "user", "content": user_message}
])
response = self.model.invoke(
[
{"role": "system", "content": system_message},
{"role": "user", "content": user_message},
]
)
# Use new .text attribute when ready
return response.content
@@ -98,7 +106,3 @@ class SummarizationMiddleware(AgentMiddleware):
def _format_messages(messages_to_summarize: Sequence[AnyMessage]) -> str:
# TODO: better formatting logic
return "\n".join([m.content for m in messages_to_summarize])
@@ -1,13 +1,84 @@
from langgraph.agent.types import AgentMiddleware, AgentState, ModelRequest
from typing import Dict, Any, List, Optional, Union
from langgraph.types import interrupt
from dataclasses import dataclass
from typing import cast
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from langchain_core.tools import BaseTool, tool
from langgraph.agent import create_agent
from langgraph.agent.types import AgentJump, AgentMiddleware, ModelRequest
class SwarmMiddleWare(AgentMiddleware):
@dataclass
class SwarmAgent:
name: str
system_prompt: str
tools: list[BaseTool]
def __init__(self, model_configs: dict[str, dict]):
super().__init__()
def modify_model_request(
self, request: ModelRequest, state: AgentState
) -> ModelRequest:
class MultiAgentMiddleware(AgentMiddleware):
"""Multi agent middleware (enabling swarm like behavior).
TODOs:
* Support create_agent for handoffs
* Support handoff customization
* do we want to include handoff messages / enable togglging
* default active agent
* handoff tool naming / descriptions
"""
@dataclass
class State(AgentMiddleware.State):
active_agent: str | None = None
@staticmethod
def _create_handoff_tools(agents: list[SwarmAgent]) -> list[BaseTool]:
handoff_tools: list[BaseTool] = []
for agent in agents:
def handoff_tool() -> str:
return f"Handing off to {agent.name}"
handoff_tools.append(
tool(
f"handoff_to_{agent.name}",
description=f"Handoff tool to trigger a handoff to {agent.name}",
)(handoff_tool)
)
return handoff_tools
def __init__(self, agents: list[SwarmAgent]):
self.agents: dict[str, SwarmAgent] = {agent.name: agent for agent in agents}
self.handoff_tools = self._create_handoff_tools(agents)
def modify_model_request(self, request: ModelRequest, state: State) -> ModelRequest:
if (active_agent := getattr(state, "active_agent", None)) is not None:
agent = self.agents[active_agent]
request.system_prompt = agent.system_prompt
request.tools = agent.tools
request.tools.extend(self.handoff_tools)
return request
def after_model(self, state) -> State | None:
# TODO: handle parallel handoffs, we don't do this currently
ai_msg: AIMessage = cast(AIMessage, state.messages[-1])
if ai_msg.tool_calls:
for call in ai_msg.tool_calls:
if call["name"].startswith("handoff_to_"):
active_agent = call["name"].replace("handoff_to_", "")
return {
"messages": [
ToolMessage(
name=call["name"],
content=f"Successfully transferred to {active_agent}",
tool_call_id=call["id"],
)
],
"active_agent": active_agent,
"jump_to": "model",
}
return None
@@ -0,0 +1,44 @@
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
class ToolCallLimitMiddleware(AgentMiddleware):
"""Terminates after a specific tool is called N times"""
@dataclass
class State(AgentMiddleware.State):
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
def after_model(self, state: State) -> AgentUpdate | AgentJump | None:
ai_msg: AIMessage = cast(AIMessage, state.messages[-1])
tool_calls = {}
for call in ai_msg.tool_calls or []:
tool_calls[call["name"]] = tool_calls.get(call["name"], 0) + 1
aggregate_calls = state.tool_call_count.copy()
for tool_name in tool_calls.keys():
aggregate_calls[tool_name] = aggregate_calls.get(tool_name, 0) + 1
for tool_name, max_calls in self.tool_limits.items():
count = aggregate_calls.get(tool_name, 0)
if count == max_calls:
return {"tool_call_count": aggregate_calls, "jump_to": "__end__"}
return {"tool_call_count": aggregate_calls}
+6 -2
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__"]
@@ -21,15 +22,16 @@ JumpTo = Literal["tools", "model", "__end__"]
class ModelRequest:
model: BaseChatModel
system_prompt: str
messages: Sequence[AnyMessage] # excluding system prompt
messages: list[AnyMessage] # excluding system prompt
tool_choice: Any
tools: Sequence[BaseTool]
tools: list[BaseTool]
response_format: ResponseFormat | None
@dataclass
class AgentState:
messages: Annotated[list[AnyMessage], add_messages]
model_request: Annotated[ModelRequest | None, EphemeralValue] = None
jump_to: Annotated[JumpTo | None, EphemeralValue] = None
response: dict | None = None
@@ -38,6 +40,8 @@ class AgentMiddleware:
class State(AgentState):
pass
tools: list[BaseTool]
def before_model(self, state: State) -> AgentUpdate | AgentJump | None:
pass
@@ -38,7 +38,6 @@ from typing_extensions import Annotated, NotRequired, TypedDict
from langgraph._internal._runnable import RunnableCallable, RunnableLike
from langgraph._internal._typing import MISSING
from langgraph.agent import create_agent
from langgraph.agent.types import AgentMiddleware
from langgraph.errors import ErrorCode, create_error_message
from langgraph.graph import END, StateGraph
@@ -472,6 +471,9 @@ def create_react_agent(
assert pre_model_hook is None
assert post_model_hook is None
assert state_schema is None
from langgraph.agent import create_agent
return create_agent(
model=model,
tools=tools,