Adding ability to skip to model, tools or END

This commit is contained in:
Nuno Campos
2025-08-26 09:57:32 +01:00
parent 80e19ecf4d
commit fdbcc07381
5 changed files with 167 additions and 18 deletions
+58 -8
View File
@@ -1,4 +1,5 @@
from collections.abc import Sequence
from inspect import signature
from typing import Callable, cast
from langchain_core.language_models.chat_models import BaseChatModel
@@ -6,9 +7,11 @@ from langchain_core.messages import AIMessage, SystemMessage
from langchain_core.tools import BaseTool
from langgraph.agent.types import (
AgentInput,
AgentGoTo,
AgentMiddleware,
AgentState,
AgentUpdate,
GoTo,
ModelRequest,
ResponseFormat,
)
@@ -24,7 +27,7 @@ def create_agent(
system_prompt: str,
middleware: Sequence[AgentMiddleware] = (),
response_format: ResponseFormat | None = None,
) -> StateGraph[AgentState, None, AgentInput]:
) -> StateGraph[AgentState, None, AgentUpdate]:
# init chat model
if isinstance(model, str):
try:
@@ -58,7 +61,7 @@ def create_agent(
]
# create graph, add nodes
graph = StateGraph(AgentState, input_schema=AgentInput)
graph = StateGraph(AgentState, input_schema=AgentUpdate, output_schema=AgentUpdate)
graph.add_node(
"model_request",
_make_model_request_node(
@@ -95,18 +98,26 @@ def create_agent(
_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])
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:]):
graph.add_edge(
_add_middleware_edge(
graph,
m1.before_model,
f"{m1.__class__.__name__}.before_model",
f"{m2.__class__.__name__}.before_model",
first_node,
)
graph.add_edge(
_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
@@ -117,9 +128,12 @@ def create_agent(
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(
_add_middleware_edge(
graph,
m1.after_model,
f"{m1.__class__.__name__}.after_model",
f"{m2.__class__.__name__}.after_model",
first_node,
)
return graph
@@ -146,6 +160,7 @@ def _make_model_request_node(
# 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
@@ -165,8 +180,17 @@ def _make_model_request_node(
return model_request
def _make_model_to_tools_edge() -> Callable[[AgentState], str | None]:
def _resolve_goto(goto: GoTo | None, first_node: str) -> str | None:
if goto == "model":
return first_node
elif goto:
return goto
def _make_model_to_tools_edge(first_node: str) -> Callable[[AgentState], str | None]:
def model_to_tools(state: AgentState) -> str | None:
if state.goto:
return _resolve_goto(state.goto, first_node)
message = state.messages[-1]
if isinstance(message, AIMessage) and message.tool_calls:
return "tools"
@@ -191,3 +215,29 @@ def _make_tools_to_model_edge(
return next_node
return tools_to_model
def _add_middleware_edge(
graph: StateGraph,
method: Callable[[AgentState], AgentUpdate | AgentGoTo | None],
name: str,
default_destination: str,
model_destination: str,
) -> None:
sig = signature(method)
uses_goto = sig.return_annotation is AgentGoTo or AgentGoTo in getattr(
sig.return_annotation, "__args__", ()
)
if uses_goto:
def goto_edge(state: AgentState) -> str:
return _resolve_goto(state.goto, model_destination) or default_destination
destinations = [default_destination, END, "tools"]
if name != model_destination:
destinations.append(model_destination)
graph.add_conditional_edges(name, goto_edge, destinations)
else:
graph.add_edge(name, default_destination)
+12 -4
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
from collections.abc import Sequence
from dataclasses import dataclass
from typing import Annotated, Any, Self
from typing import Annotated, Any, Literal, Self
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AnyMessage
@@ -10,9 +10,11 @@ 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]
GoTo = Literal["tools", "model", "__end__"]
@dataclass
@@ -32,7 +34,7 @@ class AgentMiddleware:
def __copy__(self) -> Self:
return self.__class__(**self.__dict__)
def before_model(self, state: AgentState) -> AgentState | None:
def before_model(self, state: AgentState) -> AgentUpdate | AgentGoTo | None:
pass
def modify_model_request(
@@ -40,14 +42,20 @@ class AgentMiddleware:
) -> ModelRequest:
return request
def after_model(self, state: AgentState) -> AgentState | None:
def after_model(self, state: AgentState) -> AgentUpdate | AgentGoTo | None:
pass
class AgentInput(TypedDict):
class AgentUpdate(TypedDict, total=False):
messages: Messages
class AgentGoTo(TypedDict, total=False):
messages: Messages
goto: GoTo
@dataclass
class AgentState:
messages: Annotated[list[AnyMessage], add_messages]
goto: Annotated[GoTo | None, EphemeralValue] = None
+2 -5
View File
@@ -198,11 +198,8 @@ def local_read(
# apply writes
local_channels: dict[str, BaseChannel] = {}
for k in channels:
if k in updated:
cc = channels[k].copy()
cc.update(updated[k])
else:
cc = channels[k]
cc = channels[k].copy()
cc.update(updated[k])
local_channels[k] = cc
# read fresh values
values = read_channels(local_channels, select)
@@ -277,3 +277,37 @@
'''
# ---
# name: test_create_agent_goto[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
'''
# ---
+61 -1
View File
@@ -2,8 +2,10 @@ 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.agent.types import AgentGoTo, 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
@@ -236,3 +238,61 @@ def test_create_agent_invoke(
"NoopEight.after_model",
"NoopSeven.after_model",
]
def test_create_agent_goto(
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) -> AgentGoTo:
calls.append("NoopEight.before_model")
return {"goto": 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"]