mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-18 21:55:46 +02:00
Compare commits
19
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f26ca07716 | ||
|
|
b250823532 | ||
|
|
120d34303d | ||
|
|
6ff9e4a764 | ||
|
|
3c36d2e2c8 | ||
|
|
f6d0382d66 | ||
|
|
0386fe5f6a | ||
|
|
b11ece823b | ||
|
|
46ce6ad927 | ||
|
|
0c929e62eb | ||
|
|
88c434048f | ||
|
|
f67a089a68 | ||
|
|
cc97fad7e5 | ||
|
|
75c73369a3 | ||
|
|
fdbcc07381 | ||
|
|
80e19ecf4d | ||
|
|
54272afe01 | ||
|
|
e1aeb24a4e | ||
|
|
a51c0bfa31 |
@@ -0,0 +1,314 @@
|
||||
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))
|
||||
|
||||
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), (
|
||||
"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_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
|
||||
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,
|
||||
)
|
||||
|
||||
def model_request(state: AgentState) -> AgentState:
|
||||
request = state.model_request or ModelRequest(
|
||||
model=model,
|
||||
tools=default_tools,
|
||||
system_prompt=system_prompt,
|
||||
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(
|
||||
f"{m.__class__.__name__}.before_model",
|
||||
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",
|
||||
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 f"{middleware_w_modify_model_request[0].__class__.__name__}.modify_model_request"
|
||||
if middleware_w_modify_model_request
|
||||
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), [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 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(
|
||||
"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,
|
||||
)
|
||||
# _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 _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,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,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,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}
|
||||
@@ -0,0 +1,108 @@
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from typing import Callable, Iterable
|
||||
|
||||
from langchain_core.language_models import LanguageModelLike
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AnyMessage,
|
||||
MessageLikeRepresentation,
|
||||
RemoveMessage,
|
||||
ToolMessage,
|
||||
)
|
||||
from langchain_core.messages.utils import count_tokens_approximately
|
||||
|
||||
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,84 @@
|
||||
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
|
||||
|
||||
|
||||
@dataclass
|
||||
class SwarmAgent:
|
||||
name: str
|
||||
system_prompt: str
|
||||
tools: list[BaseTool]
|
||||
|
||||
|
||||
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}
|
||||
@@ -0,0 +1,62 @@
|
||||
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
|
||||
from langgraph.runtime import Runtime
|
||||
|
||||
ResponseFormat = dict | type[BaseModel]
|
||||
JumpTo = Literal["tools", "model", "__end__"]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelRequest:
|
||||
model: BaseChatModel
|
||||
system_prompt: str
|
||||
messages: list[AnyMessage] # excluding system prompt
|
||||
tool_choice: Any
|
||||
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
|
||||
|
||||
|
||||
class AgentMiddleware:
|
||||
class State(AgentState):
|
||||
pass
|
||||
|
||||
tools: list[BaseTool]
|
||||
|
||||
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
|
||||
@@ -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
|
||||
|
||||
'''
|
||||
# ---
|
||||
@@ -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,7 @@ from typing_extensions import Annotated, NotRequired, TypedDict
|
||||
|
||||
from langgraph._internal._runnable import RunnableCallable, RunnableLike
|
||||
from langgraph._internal._typing import MISSING
|
||||
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 +262,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 +464,32 @@ 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
|
||||
|
||||
from langgraph.agent import create_agent
|
||||
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user