diff --git a/libs/langgraph/langgraph/agent/__init__.py b/libs/langgraph/langgraph/agent/__init__.py index 6ae6fe130..318cce376 100644 --- a/libs/langgraph/langgraph/agent/__init__.py +++ b/libs/langgraph/langgraph/agent/__init__.py @@ -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) diff --git a/libs/langgraph/langgraph/agent/types.py b/libs/langgraph/langgraph/agent/types.py index 2011965b0..c9fd5567d 100644 --- a/libs/langgraph/langgraph/agent/types.py +++ b/libs/langgraph/langgraph/agent/types.py @@ -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 diff --git a/libs/langgraph/langgraph/pregel/_algo.py b/libs/langgraph/langgraph/pregel/_algo.py index 2405d3d81..8eb50ffa1 100644 --- a/libs/langgraph/langgraph/pregel/_algo.py +++ b/libs/langgraph/langgraph/pregel/_algo.py @@ -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) diff --git a/libs/langgraph/tests/__snapshots__/test_agent.ambr b/libs/langgraph/tests/__snapshots__/test_agent.ambr index 8aa887354..7d380c34e 100644 --- a/libs/langgraph/tests/__snapshots__/test_agent.ambr +++ b/libs/langgraph/tests/__snapshots__/test_agent.ambr @@ -277,3 +277,37 @@ ''' # --- +# name: test_create_agent_goto[memory] + ''' + --- + config: + flowchart: + curve: linear + --- + graph TD; + __start__([

__start__

]):::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__([

__end__

]):::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 + + ''' +# --- diff --git a/libs/langgraph/tests/test_agent.py b/libs/langgraph/tests/test_agent.py index 4df2d43ad..cc66016b7 100644 --- a/libs/langgraph/tests/test_agent.py +++ b/libs/langgraph/tests/test_agent.py @@ -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"]