From fed60e713cf2bd8e76fbdce01294706853dc28a7 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 22 Nov 2024 16:28:43 -0800 Subject: [PATCH] lib: Add Command(graph=Command.PARENT, ...) - This makes the command bubble up out of the current graph and be handled by the calling graph (the immediate parent) - This could be extended to support eg. ROOT graph, or some other level --- libs/langgraph/langgraph/errors.py | 17 +++++++-- libs/langgraph/langgraph/graph/state.py | 15 +++++++- .../langgraph/langgraph/prebuilt/tool_node.py | 6 +-- libs/langgraph/langgraph/pregel/algo.py | 2 + libs/langgraph/langgraph/pregel/executor.py | 6 +-- libs/langgraph/langgraph/pregel/io.py | 3 ++ libs/langgraph/langgraph/pregel/retry.py | 38 +++++++++++++++++-- libs/langgraph/langgraph/pregel/runner.py | 6 +-- libs/langgraph/langgraph/types.py | 6 +++ 9 files changed, 82 insertions(+), 17 deletions(-) diff --git a/libs/langgraph/langgraph/errors.py b/libs/langgraph/langgraph/errors.py index 2450b42b1..0737a31d0 100644 --- a/libs/langgraph/langgraph/errors.py +++ b/libs/langgraph/langgraph/errors.py @@ -2,7 +2,7 @@ from enum import Enum from typing import Any, Sequence from langgraph.checkpoint.base import EmptyChannelError # noqa: F401 -from langgraph.types import Interrupt +from langgraph.types import Command, Interrupt # EmptyChannelError re-exported for backwards compatibility @@ -58,7 +58,11 @@ class InvalidUpdateError(Exception): pass -class GraphInterrupt(Exception): +class GraphBubbleUp(Exception): + pass + + +class GraphInterrupt(GraphBubbleUp): """Raised when a subgraph is interrupted, suppressed by the root graph. Never raised directly, or surfaced to the user.""" @@ -73,13 +77,20 @@ class NodeInterrupt(GraphInterrupt): super().__init__([Interrupt(value=value)]) -class GraphDelegate(Exception): +class GraphDelegate(GraphBubbleUp): """Raised when a graph is delegated (for distributed mode).""" def __init__(self, *args: dict[str, Any]) -> None: super().__init__(*args) +class ParentCommand(GraphBubbleUp): + args: tuple[Command] + + def __init__(self, command: Command) -> None: + super().__init__(command) + + class EmptyInputError(Exception): """Raised when graph receives an empty input.""" diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index cce742911..0684fb29c 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -37,7 +37,12 @@ from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.channels.last_value import LastValue from langgraph.channels.named_barrier_value import NamedBarrierValue from langgraph.constants import EMPTY_SEQ, NS_END, NS_SEP, SELF, TAG_HIDDEN -from langgraph.errors import ErrorCode, InvalidUpdateError, create_error_message +from langgraph.errors import ( + ErrorCode, + InvalidUpdateError, + ParentCommand, + create_error_message, +) from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph, Send from langgraph.managed.base import ( ChannelKeyPlaceholder, @@ -623,6 +628,8 @@ class CompiledStateGraph(CompiledGraph): def _get_root(input: Any) -> Any: if isinstance(input, Command): + if input.graph == Command.PARENT: + return SKIP_WRITE return input.update else: return input @@ -640,6 +647,8 @@ class CompiledStateGraph(CompiledGraph): ) return input.get(key, SKIP_WRITE) elif isinstance(input, Command): + if input.graph == Command.PARENT: + return SKIP_WRITE return _get_state_key(input.update, key=key) elif get_type_hints(type(input)): value = getattr(input, key, SKIP_WRITE) @@ -822,6 +831,8 @@ def _control_branch(value: Any) -> Sequence[Union[str, Send]]: return [value] if not isinstance(value, GraphCommand): return EMPTY_SEQ + if value.graph == Command.PARENT: + raise ParentCommand(value) rtn: list[Union[str, Send]] = [] if isinstance(value.goto, str): rtn.append(value.goto) @@ -839,6 +850,8 @@ async def _acontrol_branch(value: Any) -> Sequence[Union[str, Send]]: return [value] if not isinstance(value, GraphCommand): return EMPTY_SEQ + if value.graph == Command.PARENT: + raise ParentCommand(value) rtn: list[Union[str, Send]] = [] if isinstance(value.goto, str): rtn.append(value.goto) diff --git a/libs/langgraph/langgraph/prebuilt/tool_node.py b/libs/langgraph/langgraph/prebuilt/tool_node.py index cdcbdd819..1ea0dd56c 100644 --- a/libs/langgraph/langgraph/prebuilt/tool_node.py +++ b/libs/langgraph/langgraph/prebuilt/tool_node.py @@ -37,7 +37,7 @@ from langchain_core.tools import tool as create_tool from langchain_core.tools.base import get_all_basemodel_annotations from typing_extensions import Annotated, get_args, get_origin -from langgraph.errors import GraphInterrupt +from langgraph.errors import GraphBubbleUp from langgraph.store.base import BaseStore from langgraph.utils.runnable import RunnableCallable @@ -275,7 +275,7 @@ class ToolNode(RunnableCallable): # (2) a NodeInterrupt is raised inside a graph node for a graph called as a tool # (3) a GraphInterrupt is raised when a subgraph is interrupted inside a graph called as a tool # (2 and 3 can happen in a "supervisor w/ tools" multi-agent architecture) - except GraphInterrupt as e: + except GraphBubbleUp as e: raise e except Exception as e: if isinstance(self.handle_tool_errors, tuple): @@ -316,7 +316,7 @@ class ToolNode(RunnableCallable): # (2) a NodeInterrupt is raised inside a graph node for a graph called as a tool # (3) a GraphInterrupt is raised when a subgraph is interrupted inside a graph called as a tool # (2 and 3 can happen in a "supervisor w/ tools" multi-agent architecture) - except GraphInterrupt as e: + except GraphBubbleUp as e: raise e except Exception as e: if isinstance(self.handle_tool_errors, tuple): diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 564c53022..1410e432f 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -602,6 +602,7 @@ def prepare_single_task( None, task_id, task_path, + writers=proc.flat_writers, ) else: @@ -720,6 +721,7 @@ def prepare_single_task( None, task_id, task_path, + writers=proc.flat_writers, ) else: return PregelTask(task_id, name, task_path) diff --git a/libs/langgraph/langgraph/pregel/executor.py b/libs/langgraph/langgraph/pregel/executor.py index 246510fb4..70aea29e3 100644 --- a/libs/langgraph/langgraph/pregel/executor.py +++ b/libs/langgraph/langgraph/pregel/executor.py @@ -20,7 +20,7 @@ from langchain_core.runnables import RunnableConfig from langchain_core.runnables.config import get_executor_for_config from typing_extensions import ParamSpec -from langgraph.errors import GraphInterrupt +from langgraph.errors import GraphBubbleUp P = ParamSpec("P") T = TypeVar("T") @@ -68,7 +68,7 @@ class BackgroundExecutor(ContextManager): def done(self, task: concurrent.futures.Future) -> None: try: task.result() - except GraphInterrupt: + except GraphBubbleUp: # This exception is an interruption signal, not an error # so we don't want to re-raise it on exit self.tasks.pop(task) @@ -155,7 +155,7 @@ class AsyncBackgroundExecutor(AsyncContextManager): if exc := task.exception(): # This exception is an interruption signal, not an error # so we don't want to re-raise it on exit - if isinstance(exc, GraphInterrupt): + if isinstance(exc, GraphBubbleUp): self.tasks.pop(task) else: self.tasks.pop(task) diff --git a/libs/langgraph/langgraph/pregel/io.py b/libs/langgraph/langgraph/pregel/io.py index 6695e1ce0..693dffce2 100644 --- a/libs/langgraph/langgraph/pregel/io.py +++ b/libs/langgraph/langgraph/pregel/io.py @@ -15,6 +15,7 @@ from langgraph.constants import ( TAG_HIDDEN, TASKS, ) +from langgraph.errors import InvalidUpdateError from langgraph.pregel.log import logger from langgraph.types import Command, PregelExecutableTask, Send @@ -68,6 +69,8 @@ def map_command( cmd: Command, ) -> Iterator[tuple[str, str, Any]]: """Map input chunk to a sequence of pending writes in the form (channel, value).""" + if cmd.graph == Command.PARENT: + raise InvalidUpdateError("There is not parent graph") if cmd.send: if isinstance(cmd.send, (tuple, list)): sends = cmd.send diff --git a/libs/langgraph/langgraph/pregel/retry.py b/libs/langgraph/langgraph/pregel/retry.py index ea9162dc2..6e52a7c41 100644 --- a/libs/langgraph/langgraph/pregel/retry.py +++ b/libs/langgraph/langgraph/pregel/retry.py @@ -2,6 +2,7 @@ import asyncio import logging import random import time +from dataclasses import replace from functools import partial from typing import Any, Callable, Optional, Sequence @@ -10,9 +11,10 @@ from langgraph.constants import ( CONFIG_KEY_CHECKPOINT_NS, CONFIG_KEY_RESUMING, CONFIG_KEY_SEND, + NS_SEP, ) -from langgraph.errors import _SEEN_CHECKPOINT_NS, GraphInterrupt -from langgraph.types import PregelExecutableTask, RetryPolicy +from langgraph.errors import _SEEN_CHECKPOINT_NS, GraphBubbleUp, ParentCommand +from langgraph.types import Command, PregelExecutableTask, RetryPolicy from langgraph.utils.config import patch_configurable logger = logging.getLogger(__name__) @@ -40,7 +42,21 @@ def run_with_retry( task.proc.invoke(task.input, config) # if successful, end break - except GraphInterrupt: + except ParentCommand as exc: + ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] + cmd = exc.args[0] + if cmd.graph == ns: + # this command is for the current graph, handle it + for w in task.writers: + w.invoke(cmd, config) + break + elif cmd.graph == Command.PARENT: + # this command is for the parent graph, assign it to the parent + parent_ns = NS_SEP.join(ns.split(NS_SEP)[:-1]) + exc.args = (replace(cmd, graph=parent_ns),) + # bubble up + raise + except GraphBubbleUp: # if interrupted, end raise except Exception as exc: @@ -118,7 +134,21 @@ async def arun_with_retry( await task.proc.ainvoke(task.input, config) # if successful, end break - except GraphInterrupt: + except ParentCommand as exc: + ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS] + cmd = exc.args[0] + if cmd.graph == ns: + # this command is for the current graph, handle it + for w in task.writers: + w.invoke(cmd, config) + break + elif cmd.graph == Command.PARENT: + # this command is for the parent graph, assign it to the parent + parent_ns = NS_SEP.join(ns.split(NS_SEP)[:-1]) + exc.args = (replace(cmd, graph=parent_ns),) + # bubble up + raise + except GraphBubbleUp: # if interrupted, end raise except Exception as exc: diff --git a/libs/langgraph/langgraph/pregel/runner.py b/libs/langgraph/langgraph/pregel/runner.py index 64e5c8d3c..9e3879b0f 100644 --- a/libs/langgraph/langgraph/pregel/runner.py +++ b/libs/langgraph/langgraph/pregel/runner.py @@ -23,7 +23,7 @@ from langgraph.constants import ( PUSH, TAG_HIDDEN, ) -from langgraph.errors import GraphDelegate, GraphInterrupt +from langgraph.errors import GraphBubbleUp, GraphInterrupt from langgraph.pregel.executor import Submit from langgraph.pregel.retry import arun_with_retry, run_with_retry from langgraph.types import PregelExecutableTask, RetryPolicy @@ -298,7 +298,7 @@ class PregelRunner: # save interrupt to checkpointer if interrupts := [(INTERRUPT, i) for i in exception.args[0]]: self.put_writes(task.id, interrupts) - elif isinstance(exception, GraphDelegate): + elif isinstance(exception, GraphBubbleUp): raise exception else: # save error to checkpointer @@ -324,7 +324,7 @@ def _should_stop_others( if fut.cancelled(): return True if exc := fut.exception(): - return not isinstance(exc, GraphInterrupt) + return not isinstance(exc, GraphBubbleUp) else: return False diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 104412d8e..0a1981647 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -5,6 +5,7 @@ from typing import ( TYPE_CHECKING, Any, Callable, + ClassVar, Generic, Hashable, Literal, @@ -140,6 +141,7 @@ class PregelExecutableTask(NamedTuple): id: str path: tuple[Union[str, int, tuple], ...] scheduled: bool = False + writers: Sequence[Runnable] = () class StateSnapshot(NamedTuple): @@ -233,12 +235,14 @@ class Send: N = TypeVar("N", bound=Hashable) +PARENT = Literal["__parent__"] @dataclasses.dataclass(**_DC_KWARGS) class Command(Generic[N]): """One or more commands to update the graph's state and send messages to nodes.""" + graph: Optional[Union[PARENT, str]] = None update: Optional[dict[str, Any]] = None send: Union[Send, Sequence[Send]] = () resume: Optional[Union[Any, dict[str, Any]]] = None @@ -252,6 +256,8 @@ class Command(Generic[N]): ) return f"Command({contents})" + PARENT = ClassVar[PARENT] = "__parent__" + StreamChunk = tuple[tuple[str, ...], str, Any]