This commit is contained in:
Nuno Campos
2025-08-25 17:36:16 +01:00
parent d73902ae76
commit a51c0bfa31
4 changed files with 730 additions and 0 deletions
+178
View File
@@ -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
+47
View File
@@ -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
'''
# ---
+226
View File
@@ -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",
]