Compare commits

...
Author SHA1 Message Date
Nuno Campos 46ce6ad927 Add State property 2025-09-03 14:43:46 +01:00
Nuno Campos 0c929e62eb Rename to AgentJump 2025-09-01 09:57:52 +01:00
Nuno Campos 88c434048f Add middleware arg to create_react_agent 2025-08-29 17:07:50 +01:00
Nuno Campos f67a089a68 Rename goto to jump_to 2025-08-27 15:52:01 +01:00
Harrison Chase cc97fad7e5 cr 2025-08-26 20:26:34 -07:00
Nuno Campos 75c73369a3 Add structured response to agent output 2025-08-26 10:03:12 +01:00
Nuno Campos fdbcc07381 Adding ability to skip to model, tools or END 2025-08-26 09:57:32 +01:00
Harrison Chase 80e19ecf4d cr 2025-08-25 19:14:24 -07:00
Nuno Campos 54272afe01 Add response_format arg 2025-08-25 21:20:26 +01:00
Nuno Campos e1aeb24a4e Add state arg to modify hook 2025-08-25 21:08:38 +01:00
Nuno Campos a51c0bfa31 Boom 2025-08-25 17:36:16 +01:00
11 changed files with 1183 additions and 5 deletions
+267
View File
@@ -0,0 +1,267 @@
from collections.abc import Sequence
from inspect import signature
from typing import Any, Callable, cast
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AIMessage, SystemMessage
from langchain_core.tools import BaseTool
from langgraph.agent.types import (
AgentJump,
AgentMiddleware,
AgentState,
AgentUpdate,
JumpTo,
ModelRequest,
ResponseFormat,
)
from langgraph.constants import END, START
from langgraph.graph.state import StateGraph
from langgraph.prebuilt.tool_node import ToolNode
def create_agent(
*,
model: str | BaseChatModel,
tools: Sequence[BaseTool | Callable] | ToolNode,
system_prompt: str,
middleware: Sequence[AgentMiddleware] = (),
response_format: ResponseFormat | None = None,
context_schema: type[Any] | None = None,
) -> StateGraph[AgentState, None, AgentUpdate]:
# init chat model
if isinstance(model, str):
try:
from langchain.chat_models import ( # type: ignore[import-not-found]
init_chat_model,
)
except ImportError:
raise ImportError(
"Please install langchain (`pip install langchain`) to "
"use '<provider>:<model>' string syntax for `model` parameter."
)
model = cast(BaseChatModel, init_chat_model(model))
# init tool node
tool_node = tools if isinstance(tools, ToolNode) else ToolNode(tools=tools)
# validate middleware
assert len({m.__class__.__name__ for m in middleware}) == len(middleware), (
"Please remove duplicate middleware instances."
) # this is just to keep the node names simple, we can change if needed
middleware_w_before = [
m
for m in middleware
if m.__class__.before_model is not AgentMiddleware.before_model
]
middleware_w_after = [
m
for m in middleware
if m.__class__.after_model is not AgentMiddleware.after_model
]
# create graph, add nodes
graph = StateGraph(
AgentState,
input_schema=AgentUpdate,
output_schema=AgentUpdate,
context_schema=context_schema,
)
graph.add_node(
"model_request",
_make_model_request_node(
model=model,
tools=list(tool_node.tools_by_name.values()),
system_prompt=system_prompt,
middleware=middleware,
response_format=response_format,
),
)
graph.add_node("tools", tool_node)
for m in middleware:
if m.__class__.before_model is not AgentMiddleware.before_model:
graph.add_node(
f"{m.__class__.__name__}.before_model",
m.before_model,
input_schema=m.State,
)
if m.__class__.after_model is not AgentMiddleware.after_model:
graph.add_node(
f"{m.__class__.__name__}.after_model",
m.after_model,
input_schema=m.State,
)
# add start edge
first_node = (
f"{middleware_w_before[0].__class__.__name__}.before_model"
if middleware_w_before
else "model_request"
)
last_node = (
f"{middleware_w_after[0].__class__.__name__}.after_model"
if middleware_w_after
else "model_request"
)
graph.add_edge(START, first_node)
# add cond edges
graph.add_conditional_edges(
"tools",
_make_tools_to_model_edge(tool_node, first_node),
[first_node, END],
)
graph.add_conditional_edges(
last_node, _make_model_to_tools_edge(first_node), ["tools", END]
)
# add before model edges
if middleware_w_before:
for m1, m2 in zip(middleware_w_before, middleware_w_before[1:]):
_add_middleware_edge(
graph,
m1.before_model,
f"{m1.__class__.__name__}.before_model",
f"{m2.__class__.__name__}.before_model",
first_node,
)
_add_middleware_edge(
graph,
middleware_w_before[-1].before_model,
f"{middleware_w_before[-1].__class__.__name__}.before_model",
"model_request",
first_node,
)
# add after model edges
if middleware_w_after:
graph.add_edge(
"model_request", f"{middleware_w_after[-1].__class__.__name__}.after_model"
)
for idx in range(len(middleware_w_after) - 1, 0, -1):
m1 = middleware_w_after[idx]
m2 = middleware_w_after[idx - 1]
_add_middleware_edge(
graph,
m1.after_model,
f"{m1.__class__.__name__}.after_model",
f"{m2.__class__.__name__}.after_model",
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
elif jump_to:
return jump_to
def _make_model_to_tools_edge(first_node: str) -> Callable[[AgentState], str | None]:
def model_to_tools(state: AgentState) -> str | None:
if state.jump_to:
return _resolve_jump(state.jump_to, first_node)
message = state.messages[-1]
if isinstance(message, AIMessage) and message.tool_calls:
return "tools"
return END
return model_to_tools
def _make_tools_to_model_edge(
tool_node: ToolNode, next_node: str
) -> Callable[[AgentState], str | None]:
def tools_to_model(state: AgentState) -> str | None:
ai_message = [m for m in state.messages if isinstance(m, AIMessage)][-1]
if all(
tool_node.tools_by_name[c["name"]].return_direct
for c in ai_message.tool_calls
if c["name"] in tool_node.tools_by_name
):
return END
return next_node
return tools_to_model
def _add_middleware_edge(
graph: StateGraph,
method: Callable[[AgentState], AgentUpdate | AgentJump | None],
name: str,
default_destination: str,
model_destination: str,
) -> None:
sig = signature(method)
uses_jump = sig.return_annotation is AgentJump or AgentJump in getattr(
sig.return_annotation, "__args__", ()
)
if uses_jump:
def jump_edge(state: AgentState) -> str:
return (
_resolve_jump(state.jump_to, model_destination) or default_destination
)
destinations = [default_destination, END, "tools"]
if name != model_destination:
destinations.append(model_destination)
graph.add_conditional_edges(name, jump_edge, destinations)
else:
graph.add_edge(name, default_destination)
@@ -0,0 +1,102 @@
from langgraph.agent.types import AgentJump, AgentMiddleware, AgentState, AgentUpdate
from langgraph.prebuilt.interrupt import (
ActionRequest,
HumanInterrupt,
HumanInterruptConfig,
HumanResponse,
)
from langgraph.types import interrupt
ToolInterruptConfig = dict[str, HumanInterruptConfig]
class HumanInTheLoopMiddleware(AgentMiddleware):
def __init__(
self,
tool_configs: ToolInterruptConfig,
message_prefix: str = "Tool execution requires approval",
):
super().__init__()
self.tool_configs = tool_configs
self.message_prefix = message_prefix
def after_model(self, state: AgentState) -> AgentUpdate | AgentJump | None:
messages = state.messages
if not messages:
return
last_message = messages[-1]
if not hasattr(last_message, "tool_calls") or not last_message.tool_calls:
return
# Separate tool calls that need interrupts from those that don't
interrupt_tool_calls = []
auto_approved_tool_calls = []
for tool_call in last_message.tool_calls:
tool_name = tool_call["name"]
if tool_name in self.tool_configs:
interrupt_tool_calls.append(tool_call)
else:
auto_approved_tool_calls.append(tool_call)
# If no interrupts needed, return early
if not interrupt_tool_calls:
return
approved_tool_calls = auto_approved_tool_calls.copy()
# Process all tool calls that need interrupts in parallel
requests = []
for tool_call in interrupt_tool_calls:
tool_name = tool_call["name"]
tool_args = tool_call["args"]
description = (
f"{self.message_prefix}\n\nTool: {tool_name}\nArgs: {tool_args}"
)
tool_config = self.tool_configs[tool_name]
request: HumanInterrupt = {
"action_request": ActionRequest(
action=tool_name,
args=tool_args,
),
"config": tool_config,
"description": description,
}
requests.append(request)
responses: list[HumanResponse] = interrupt(requests)
for i, response in enumerate(responses):
tool_call = interrupt_tool_calls[i]
if response["type"] == "accept":
approved_tool_calls.append(tool_call)
elif response["type"] == "edit":
edited: ActionRequest = response["args"]
new_tool_call = {
"name": tool_call["name"],
"args": edited["args"],
"id": tool_call["id"],
}
approved_tool_calls.append(new_tool_call)
elif response["type"] == "ignore":
# NOTE: does not work with multiple interrupts
return {"goto": "__end__"}
elif response["type"] == "response":
# NOTE: does not work with multiple interrupts
tool_message = {
"role": "tool",
"tool_call_id": tool_call["id"],
"content": response["args"],
}
return {"messages": [tool_message], "goto": "model"}
else:
raise ValueError(f"Unknown response type: {response['type']}")
last_message.tool_calls = approved_tool_calls
return {"messages": [last_message]}
@@ -0,0 +1,104 @@
from typing import Callable, Iterable
from langchain_core.language_models import LanguageModelLike
from langchain_core.messages import RemoveMessage, MessageLikeRepresentation
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>
<primary_objective>
Your sole objective in this task is to extract the highest quality/most relevant context from the conversation history below.
</primary_objective>
<objective_information>
You're nearing the total number of input tokens you can accept, so you must extract the highest quality/most relevant pieces of information from your conversation history.
This context will then overwrite the conversation history presented below. Because of this, ensure the context you extract is only the most important information to your overall goal.
</objective_information>
<instructions>
The conversation history below will be replaced with the context you extract in this step. Because of this, you must do your very best to extract and record all of the most important context from the conversation history.
You want to ensure that you don't repeat any actions you've already completed, so the context you extract from the conversation history should be focused on the most important information to your overall goal.
</instructions>
The user will message you with the full message history you'll be extracting context from, to then replace. Carefully read over it all, and think deeply about what information is most important to your overall goal that should be saved:
With all of this in mind, please carefully read over the entire conversation history, and extract the most important and relevant context to replace it so that you can free up space in the conversation history.
Respond ONLY with the extracted context. Do not include any additional information, or text before or after the extracted context."""
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"):
super().__init__()
self.model = model
self.max_tokens_before_summary = max_tokens_before_summary
self.token_counter = token_counter
self.messages_to_leave = messages_to_leave
self.summary_system_prompt = summary_system_prompt
self.fake_tool_call_name = fake_tool_call_name
def before_model(self, state: AgentState) -> AgentState | None:
messages = state.messages
token_counts = self.token_counter(messages)
# If token counts are less than max allowed, then end this hook early
if token_counts < self.max_tokens_before_summary:
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]
# 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)]
return {
"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}
])
# Use new .text attribute when ready
return response.content
@staticmethod
def _format_messages(messages_to_summarize: Sequence[AnyMessage]) -> str:
# TODO: better formatting logic
return "\n".join([m.content for m in messages_to_summarize])
@@ -0,0 +1,13 @@
from langgraph.agent.types import AgentMiddleware, AgentState, ModelRequest
from typing import Dict, Any, List, Optional, Union
from langgraph.types import interrupt
class SwarmMiddleWare(AgentMiddleware):
def __init__(self, model_configs: dict[str, dict]):
super().__init__()
def modify_model_request(
self, request: ModelRequest, state: AgentState
) -> ModelRequest:
+58
View File
@@ -0,0 +1,58 @@
from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
from typing import Annotated, Any, Literal
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AnyMessage
from langchain_core.tools import BaseTool
from pydantic import BaseModel
from typing_extensions import TypedDict
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.graph.message import Messages, add_messages
ResponseFormat = dict | type[BaseModel]
JumpTo = Literal["tools", "model", "__end__"]
@dataclass
class ModelRequest:
model: BaseChatModel
system_prompt: str
messages: Sequence[AnyMessage] # excluding system prompt
tool_choice: Any
tools: Sequence[BaseTool]
response_format: ResponseFormat | None
@dataclass
class AgentState:
messages: Annotated[list[AnyMessage], add_messages]
jump_to: Annotated[JumpTo | None, EphemeralValue] = None
response: dict | None = None
class AgentMiddleware:
class State(AgentState):
pass
def before_model(self, state: State) -> AgentUpdate | AgentJump | None:
pass
def modify_model_request(self, request: ModelRequest, state: State) -> ModelRequest:
return request
def after_model(self, state: State) -> AgentUpdate | AgentJump | None:
pass
class AgentUpdate(TypedDict, total=False):
messages: Messages
response: dict
class AgentJump(TypedDict, total=False):
messages: Messages
jump_to: JumpTo
+2 -5
View File
@@ -198,11 +198,8 @@ def local_read(
# apply writes
local_channels: dict[str, BaseChannel] = {}
for k in channels:
if k in updated:
cc = channels[k].copy()
cc.update(updated[k])
else:
cc = channels[k]
cc = channels[k].copy()
cc.update(updated[k])
local_channels[k] = cc
# read fresh values
values = read_channels(local_channels, select)
@@ -0,0 +1,313 @@
# serializer version: 1
# name: test_create_agent_diagram
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
__end__([<p>__end__</p>]):::last
__start__ --> model_request;
model_request -.-> __end__;
model_request -.-> tools;
tools -.-> __end__;
tools -.-> model_request;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.1
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopOne_before_model(NoopOne.before_model)
__end__([<p>__end__</p>]):::last
NoopOne_before_model --> model_request;
__start__ --> NoopOne_before_model;
model_request -.-> __end__;
model_request -.-> tools;
tools -.-> NoopOne_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.2
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopOne_before_model(NoopOne.before_model)
NoopTwo_before_model(NoopTwo.before_model)
__end__([<p>__end__</p>]):::last
NoopOne_before_model --> NoopTwo_before_model;
NoopTwo_before_model --> model_request;
__start__ --> NoopOne_before_model;
model_request -.-> __end__;
model_request -.-> tools;
tools -.-> NoopOne_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.3
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopOne_before_model(NoopOne.before_model)
NoopTwo_before_model(NoopTwo.before_model)
NoopThree_before_model(NoopThree.before_model)
__end__([<p>__end__</p>]):::last
NoopOne_before_model --> NoopTwo_before_model;
NoopThree_before_model --> model_request;
NoopTwo_before_model --> NoopThree_before_model;
__start__ --> NoopOne_before_model;
model_request -.-> __end__;
model_request -.-> tools;
tools -.-> NoopOne_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.4
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopFour_after_model(NoopFour.after_model)
__end__([<p>__end__</p>]):::last
NoopFour_after_model -.-> __end__;
NoopFour_after_model -.-> tools;
__start__ --> model_request;
model_request --> NoopFour_after_model;
tools -.-> __end__;
tools -.-> model_request;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.5
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopFour_after_model(NoopFour.after_model)
NoopFive_after_model(NoopFive.after_model)
__end__([<p>__end__</p>]):::last
NoopFive_after_model --> NoopFour_after_model;
NoopFour_after_model -.-> __end__;
NoopFour_after_model -.-> tools;
__start__ --> model_request;
model_request --> NoopFive_after_model;
tools -.-> __end__;
tools -.-> model_request;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.6
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopFour_after_model(NoopFour.after_model)
NoopFive_after_model(NoopFive.after_model)
NoopSix_after_model(NoopSix.after_model)
__end__([<p>__end__</p>]):::last
NoopFive_after_model --> NoopFour_after_model;
NoopFour_after_model -.-> __end__;
NoopFour_after_model -.-> tools;
NoopSix_after_model --> NoopFive_after_model;
__start__ --> model_request;
model_request --> NoopSix_after_model;
tools -.-> __end__;
tools -.-> model_request;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.7
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopSeven_before_model(NoopSeven.before_model)
NoopSeven_after_model(NoopSeven.after_model)
__end__([<p>__end__</p>]):::last
NoopSeven_after_model -.-> __end__;
NoopSeven_after_model -.-> tools;
NoopSeven_before_model --> model_request;
__start__ --> NoopSeven_before_model;
model_request --> NoopSeven_after_model;
tools -.-> NoopSeven_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.8
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopSeven_before_model(NoopSeven.before_model)
NoopSeven_after_model(NoopSeven.after_model)
NoopEight_before_model(NoopEight.before_model)
NoopEight_after_model(NoopEight.after_model)
__end__([<p>__end__</p>]):::last
NoopEight_after_model --> NoopSeven_after_model;
NoopEight_before_model --> model_request;
NoopSeven_after_model -.-> __end__;
NoopSeven_after_model -.-> tools;
NoopSeven_before_model --> NoopEight_before_model;
__start__ --> NoopSeven_before_model;
model_request --> NoopEight_after_model;
tools -.-> NoopSeven_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_diagram.9
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopSeven_before_model(NoopSeven.before_model)
NoopSeven_after_model(NoopSeven.after_model)
NoopEight_before_model(NoopEight.before_model)
NoopEight_after_model(NoopEight.after_model)
NoopNine_before_model(NoopNine.before_model)
NoopNine_after_model(NoopNine.after_model)
__end__([<p>__end__</p>]):::last
NoopEight_after_model --> NoopSeven_after_model;
NoopEight_before_model --> NoopNine_before_model;
NoopNine_after_model --> NoopEight_after_model;
NoopNine_before_model --> model_request;
NoopSeven_after_model -.-> __end__;
NoopSeven_after_model -.-> tools;
NoopSeven_before_model --> NoopEight_before_model;
__start__ --> NoopSeven_before_model;
model_request --> NoopNine_after_model;
tools -.-> NoopSeven_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_create_agent_jump[memory]
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
model_request(model_request)
tools(tools)
NoopSeven_before_model(NoopSeven.before_model)
NoopSeven_after_model(NoopSeven.after_model)
NoopEight_before_model(NoopEight.before_model)
NoopEight_after_model(NoopEight.after_model)
__end__([<p>__end__</p>]):::last
NoopEight_after_model --> NoopSeven_after_model;
NoopEight_before_model -.-> NoopSeven_before_model;
NoopEight_before_model -.-> __end__;
NoopEight_before_model -.-> model_request;
NoopEight_before_model -.-> tools;
NoopSeven_after_model -.-> __end__;
NoopSeven_after_model -.-> tools;
NoopSeven_before_model --> NoopEight_before_model;
__start__ --> NoopSeven_before_model;
model_request --> NoopEight_after_model;
tools -.-> NoopSeven_before_model;
tools -.-> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
+298
View File
@@ -0,0 +1,298 @@
from langchain_core.messages import AIMessage, ToolCall
from syrupy import SnapshotAssertion
from langgraph.agent import create_agent
from langgraph.agent.types import AgentJump, AgentMiddleware
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.constants import END
from tests.fake_chat import FakeChatModel
from tests.messages import _AnyIdToolMessage
def test_create_agent_diagram(
snapshot: SnapshotAssertion,
):
class NoopOne(AgentMiddleware):
def before_model(self, state):
pass
class NoopTwo(AgentMiddleware):
def before_model(self, state):
pass
class NoopThree(AgentMiddleware):
def before_model(self, state):
pass
class NoopFour(AgentMiddleware):
def after_model(self, state):
pass
class NoopFive(AgentMiddleware):
def after_model(self, state):
pass
class NoopSix(AgentMiddleware):
def after_model(self, state):
pass
class NoopSeven(AgentMiddleware):
def before_model(self, state):
pass
def after_model(self, state):
pass
class NoopEight(AgentMiddleware):
def before_model(self, state):
pass
def after_model(self, state):
pass
class NoopNine(AgentMiddleware):
def before_model(self, state):
pass
def after_model(self, state):
pass
agent_zero = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
)
assert agent_zero.compile().get_graph().draw_mermaid() == snapshot
agent_one = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopOne()],
)
assert agent_one.compile().get_graph().draw_mermaid() == snapshot
agent_two = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopOne(), NoopTwo()],
)
assert agent_two.compile().get_graph().draw_mermaid() == snapshot
agent_three = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopOne(), NoopTwo(), NoopThree()],
)
assert agent_three.compile().get_graph().draw_mermaid() == snapshot
agent_four = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopFour()],
)
assert agent_four.compile().get_graph().draw_mermaid() == snapshot
agent_five = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopFour(), NoopFive()],
)
assert agent_five.compile().get_graph().draw_mermaid() == snapshot
agent_six = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopFour(), NoopFive(), NoopSix()],
)
assert agent_six.compile().get_graph().draw_mermaid() == snapshot
agent_seven = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopSeven()],
)
assert agent_seven.compile().get_graph().draw_mermaid() == snapshot
agent_eight = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopSeven(), NoopEight()],
)
assert agent_eight.compile().get_graph().draw_mermaid() == snapshot
agent_nine = create_agent(
model=FakeChatModel(messages=[]),
tools=[],
system_prompt="You are a helpful assistant.",
middleware=[NoopSeven(), NoopEight(), NoopNine()],
)
assert agent_nine.compile().get_graph().draw_mermaid() == snapshot
def test_create_agent_invoke(
snapshot: SnapshotAssertion,
sync_checkpointer: BaseCheckpointSaver,
):
calls = []
class NoopSeven(AgentMiddleware):
def before_model(self, state):
calls.append("NoopSeven.before_model")
def modify_model_request(self, request, state):
calls.append("NoopSeven.modify_model_request")
return request
def after_model(self, state):
calls.append("NoopSeven.after_model")
class NoopEight(AgentMiddleware):
def before_model(self, state):
calls.append("NoopEight.before_model")
def modify_model_request(self, request, state):
calls.append("NoopEight.modify_model_request")
return request
def after_model(self, state):
calls.append("NoopEight.after_model")
def my_tool(input: str) -> str:
"""A great tool"""
calls.append("my_tool")
return input.upper()
agent_one = create_agent(
model=FakeChatModel(
messages=[
AIMessage(
"",
id="ai1",
tool_calls=[ToolCall(id="1", name="my_tool", args={"input": "yo"})],
),
AIMessage(id="ai2", content="Hello, how can I assist you today?"),
]
),
tools=[my_tool],
system_prompt="You are a helpful assistant.",
middleware=[NoopSeven(), NoopEight()],
).compile(checkpointer=sync_checkpointer)
thread1 = {"configurable": {"thread_id": "1"}}
assert agent_one.invoke({"messages": []}, thread1) == {
"messages": [
AIMessage(
content="",
additional_kwargs={},
response_metadata={},
id="ai1",
tool_calls=[
{
"name": "my_tool",
"args": {"input": "yo"},
"id": "1",
"type": "tool_call",
}
],
),
_AnyIdToolMessage(content="YO", name="my_tool", tool_call_id="1"),
AIMessage(
content="Hello, how can I assist you today?",
additional_kwargs={},
response_metadata={},
id="ai2",
),
],
}
assert calls == [
"NoopSeven.before_model",
"NoopEight.before_model",
"NoopSeven.modify_model_request",
"NoopEight.modify_model_request",
"NoopEight.after_model",
"NoopSeven.after_model",
"my_tool",
"NoopSeven.before_model",
"NoopEight.before_model",
"NoopSeven.modify_model_request",
"NoopEight.modify_model_request",
"NoopEight.after_model",
"NoopSeven.after_model",
]
def test_create_agent_jump(
snapshot: SnapshotAssertion,
sync_checkpointer: BaseCheckpointSaver,
):
calls = []
class NoopSeven(AgentMiddleware):
def before_model(self, state):
calls.append("NoopSeven.before_model")
def modify_model_request(self, request, state):
calls.append("NoopSeven.modify_model_request")
return request
def after_model(self, state):
calls.append("NoopSeven.after_model")
class NoopEight(AgentMiddleware):
def before_model(self, state) -> AgentJump:
calls.append("NoopEight.before_model")
return {"jump_to": END}
def modify_model_request(self, request, state):
calls.append("NoopEight.modify_model_request")
return request
def after_model(self, state):
calls.append("NoopEight.after_model")
def my_tool(input: str) -> str:
"""A great tool"""
calls.append("my_tool")
return input.upper()
agent_one = create_agent(
model=FakeChatModel(
messages=[
AIMessage(
"",
id="ai1",
tool_calls=[ToolCall(id="1", name="my_tool", args={"input": "yo"})],
),
AIMessage(id="ai2", content="Hello, how can I assist you today?"),
]
),
tools=[my_tool],
system_prompt="You are a helpful assistant.",
middleware=[NoopSeven(), NoopEight()],
).compile(checkpointer=sync_checkpointer)
if isinstance(sync_checkpointer, InMemorySaver):
assert agent_one.get_graph().draw_mermaid() == snapshot
thread1 = {"configurable": {"thread_id": "1"}}
assert agent_one.invoke({"messages": []}, thread1) == {"messages": []}
assert calls == ["NoopSeven.before_model", "NoopEight.before_model"]
@@ -38,6 +38,8 @@ 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
from langgraph.graph.message import add_messages
@@ -261,6 +263,7 @@ def create_react_agent(
],
tools: Union[Sequence[Union[BaseTool, Callable, dict[str, Any]]], ToolNode],
*,
middleware: Sequence[AgentMiddleware] = (),
prompt: Optional[Prompt] = None,
response_format: Optional[
Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]]
@@ -462,6 +465,29 @@ def create_react_agent(
print(chunk)
```
"""
if middleware:
assert isinstance(model, str | BaseChatModel)
assert isinstance(prompt, str | None)
assert not isinstance(response_format, tuple)
assert pre_model_hook is None
assert post_model_hook is None
assert state_schema is None
return create_agent(
model=model,
tools=tools,
system_prompt=prompt,
middleware=middleware,
response_format=response_format,
context_schema=context_schema,
).compile(
checkpointer=checkpointer,
store=store,
name=name,
interrupt_after=interrupt_after,
interrupt_before=interrupt_before,
debug=debug,
)
if (
config_schema := deprecated_kwargs.pop("config_schema", MISSING)
) is not MISSING: