mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-30 11:49:38 +02:00
Boom
This commit is contained in:
@@ -0,0 +1,178 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import 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 AgentInput, AgentMiddleware, AgentState, ModelRequest
|
||||
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],
|
||||
system_prompt: str,
|
||||
middleware: Sequence[AgentMiddleware] = (),
|
||||
) -> StateGraph[AgentState, None, AgentInput]:
|
||||
# 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 = 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=AgentInput)
|
||||
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,
|
||||
),
|
||||
)
|
||||
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)
|
||||
if m.__class__.after_model is not AgentMiddleware.after_model:
|
||||
graph.add_node(f"{m.__class__.__name__}.after_model", m.after_model)
|
||||
|
||||
# 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(), ["tools", END])
|
||||
|
||||
# add before model edges
|
||||
if middleware_w_before:
|
||||
for m1, m2 in zip(middleware_w_before, middleware_w_before[1:]):
|
||||
graph.add_edge(
|
||||
f"{m1.__class__.__name__}.before_model",
|
||||
f"{m2.__class__.__name__}.before_model",
|
||||
)
|
||||
graph.add_edge(
|
||||
f"{middleware_w_before[-1].__class__.__name__}.before_model",
|
||||
"model_request",
|
||||
)
|
||||
|
||||
# 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]
|
||||
graph.add_edge(
|
||||
f"{m1.__class__.__name__}.after_model",
|
||||
f"{m2.__class__.__name__}.after_model",
|
||||
)
|
||||
|
||||
return graph
|
||||
|
||||
|
||||
def _make_model_request_node(
|
||||
*,
|
||||
system_prompt: str,
|
||||
model: BaseChatModel,
|
||||
tools: Sequence[BaseTool],
|
||||
middleware: Sequence[AgentMiddleware],
|
||||
) -> 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,
|
||||
)
|
||||
# visit middleware in order
|
||||
for mw in middleware:
|
||||
request = mw.modify_model_request(request)
|
||||
# prepare messages
|
||||
if request.system_prompt:
|
||||
messages = [SystemMessage(request.system_prompt)] + request.messages
|
||||
else:
|
||||
messages = request.messages
|
||||
# call model
|
||||
output = request.model.invoke(
|
||||
messages, tools=request.tools, tool_choice=request.tool_choice
|
||||
)
|
||||
return {"messages": output}
|
||||
|
||||
return model_request
|
||||
|
||||
|
||||
def _make_model_to_tools_edge() -> Callable[[AgentState], str | None]:
|
||||
def model_to_tools(state: AgentState) -> str | None:
|
||||
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
|
||||
@@ -0,0 +1,47 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Annotated, Any, Self
|
||||
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import AnyMessage
|
||||
from langchain_core.tools import BaseTool
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.graph.message import Messages, add_messages
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelRequest:
|
||||
model: BaseChatModel
|
||||
system_prompt: str
|
||||
messages: Sequence[AnyMessage] # excluding system prompt
|
||||
tool_choice: Any
|
||||
tools: Sequence[BaseTool]
|
||||
|
||||
|
||||
class AgentMiddleware:
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def __copy__(self) -> Self:
|
||||
return self.__class__(**self.__dict__)
|
||||
|
||||
def before_model(self, state: AgentState) -> AgentState | None:
|
||||
pass
|
||||
|
||||
def modify_model_request(self, request: ModelRequest) -> ModelRequest:
|
||||
return request
|
||||
|
||||
def after_model(self, state: AgentState) -> AgentState | None:
|
||||
pass
|
||||
|
||||
|
||||
class AgentInput(TypedDict):
|
||||
messages: Messages
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentState:
|
||||
messages: Annotated[list[AnyMessage], add_messages]
|
||||
@@ -0,0 +1,279 @@
|
||||
# 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
|
||||
|
||||
'''
|
||||
# ---
|
||||
@@ -0,0 +1,226 @@
|
||||
from langchain_core.messages import AIMessage, ToolCall
|
||||
from syrupy import SnapshotAssertion
|
||||
|
||||
from langgraph.agent import create_agent
|
||||
from langgraph.agent.types import AgentMiddleware
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
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 after_model(self, state):
|
||||
calls.append("NoopSeven.after_model")
|
||||
|
||||
class NoopEight(AgentMiddleware):
|
||||
def before_model(self, state):
|
||||
calls.append("NoopEight.before_model")
|
||||
|
||||
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",
|
||||
"NoopEight.after_model",
|
||||
"NoopSeven.after_model",
|
||||
"my_tool",
|
||||
"NoopSeven.before_model",
|
||||
"NoopEight.before_model",
|
||||
"NoopEight.after_model",
|
||||
"NoopSeven.after_model",
|
||||
]
|
||||
Reference in New Issue
Block a user