mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
Compare commits
11
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
46ce6ad927 | ||
|
|
0c929e62eb | ||
|
|
88c434048f | ||
|
|
f67a089a68 | ||
|
|
cc97fad7e5 | ||
|
|
75c73369a3 | ||
|
|
fdbcc07381 | ||
|
|
80e19ecf4d | ||
|
|
54272afe01 | ||
|
|
e1aeb24a4e | ||
|
|
a51c0bfa31 |
@@ -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:
|
||||
@@ -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
|
||||
@@ -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,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:
|
||||
|
||||
Reference in New Issue
Block a user