mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-22 23:52:23 +02:00
Adding ability to skip to model, tools or END
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
'''
|
||||
# ---
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user