From 64086aa814939c11723767c9bec6e1c399ccfd25 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 10 Apr 2025 17:25:50 -0700 Subject: [PATCH] Avoid validating node input more than once per superstep --- libs/langgraph/langgraph/graph/state.py | 26 +++++++-- libs/langgraph/langgraph/pregel/algo.py | 71 ++++++++++++++----------- libs/langgraph/langgraph/pregel/read.py | 13 +++++ 3 files changed, 75 insertions(+), 35 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 093faf5d4..9555be1be 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -638,6 +638,7 @@ class StateGraph(Graph): compiled = CompiledStateGraph( builder=self, + schema_to_mapper={}, config_type=self.config_schema, input_model=( self.input @@ -688,6 +689,16 @@ class StateGraph(Graph): class CompiledStateGraph(CompiledGraph): builder: StateGraph + schema_to_mapper: dict[Type[Any], Optional[Callable[[Any], Any]]] + + def __init__( + self, + *, + schema_to_mapper: dict[Type[Any], Optional[Callable[[Any], Any]]], + **kwargs: Any, + ) -> None: + super().__init__(**kwargs) + self.schema_to_mapper = schema_to_mapper def get_input_schema( self, config: Optional[RunnableConfig] = None @@ -827,6 +838,15 @@ class CompiledStateGraph(CompiledGraph): input_schema = node.input if node else self.builder.schema input_values = {k: k for k in self.builder.schemas[input_schema]} is_single_input = len(input_values) == 1 and "__root__" in input_values + if input_schema in self.schema_to_mapper: + mapper = self.schema_to_mapper[input_schema] + else: + mapper = _pick_mapper( + list(input_values), + input_schema, + self.builder.type_hints[input_schema], + ) + self.schema_to_mapper[input_schema] = mapper branch_channel = CHANNEL_BRANCH_TO.format(key) self.channels[branch_channel] = EphemeralValue(Any, guard=False) @@ -835,11 +855,7 @@ class CompiledStateGraph(CompiledGraph): # read state keys and managed values channels=(list(input_values) if is_single_input else input_values), # coerce state dict to schema class (eg. pydantic model) - mapper=_pick_mapper( - list(input_values), - input_schema, - self.builder.type_hints[input_schema], - ), + mapper=mapper, # publish to state keys writers=[ChannelWrite(write_entries, tags=[TAG_HIDDEN])], metadata=node.metadata, diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index e43890498..d6bcaf6ea 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -3,13 +3,13 @@ import itertools import sys import threading from collections import defaultdict, deque +from copy import copy from functools import partial from hashlib import sha1 from typing import ( Any, Callable, Iterable, - Iterator, Literal, Mapping, NamedTuple, @@ -49,6 +49,7 @@ from langgraph.constants import ( EMPTY_SEQ, ERROR, INTERRUPT, + MISSING, NO_WRITES, NS_END, NS_SEP, @@ -63,12 +64,12 @@ from langgraph.constants import ( TASKS, Send, ) -from langgraph.errors import EmptyChannelError, InvalidUpdateError +from langgraph.errors import InvalidUpdateError from langgraph.managed.base import ManagedValueMapping from langgraph.pregel.call import get_runnable_for_task -from langgraph.pregel.io import read_channel, read_channels +from langgraph.pregel.io import read_channels from langgraph.pregel.log import logger -from langgraph.pregel.read import PregelNode +from langgraph.pregel.read import INPUT_CACHE_KEY_TYPE, PregelNode from langgraph.store.base import BaseStore from langgraph.types import ( All, @@ -423,6 +424,7 @@ def prepare_next_tasks( are the tasks themselves. This is the union of all PUSH tasks (Sends) and PULL tasks (nodes triggered by edges). """ + input_cache: dict[tuple[Callable[..., Any]], tuple[str, ...]] = {} checkpoint_id_bytes = binascii.unhexlify(checkpoint["id"].replace("-", "")) null_version = checkpoint_null_version(checkpoint) tasks: list[Union[PregelTask, PregelExecutableTask]] = [] @@ -444,6 +446,7 @@ def prepare_next_tasks( store=store, checkpointer=checkpointer, manager=manager, + input_cache=input_cache, ): tasks.append(task) @@ -486,6 +489,7 @@ def prepare_next_tasks( store=store, checkpointer=checkpointer, manager=manager, + input_cache=input_cache, ): tasks.append(task) return {t.id: t for t in tasks} @@ -511,6 +515,7 @@ def prepare_single_task( store: Optional[BaseStore] = None, checkpointer: Optional[BaseCheckpointSaver] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, + input_cache: Optional[dict[tuple[Callable[..., Any]], tuple[str, ...]]] = None, ) -> Union[None, PregelTask, PregelExecutableTask]: """Prepares a single task for the next Pregel step, given a task path, which uniquely identifies a PUSH or PULL task within the graph.""" @@ -729,11 +734,15 @@ def prepare_single_task( ): triggers = tuple(sorted(proc.triggers)) try: - val = next( - _proc_input(proc, managed, channels, for_execution=for_execution) + val = _proc_input( + proc, + managed, + channels, + for_execution=for_execution, + input_cache=input_cache, ) - except StopIteration: - return + if val is MISSING: + return except Exception as exc: if SUPPORTS_EXC_NOTES: exc.add_note( @@ -926,34 +935,32 @@ def _proc_input( channels: Mapping[str, BaseChannel], *, for_execution: bool, -) -> Iterator[Any]: + input_cache: Optional[dict[INPUT_CACHE_KEY_TYPE, Any]], +) -> Any: """Prepare input for a PULL task, based on the process's channels and triggers.""" + # if in cache return shallow copy + if input_cache is not None and proc.input_cache_key in input_cache: + return copy(input_cache[proc.input_cache_key]) # If all trigger channels subscribed by this process are not empty # then invoke the process with the values of all non-empty channels if isinstance(proc.channels, dict): - try: - val: dict[str, Any] = {} - for k, chan in proc.channels.items(): - if chan in proc.triggers: - val[k] = read_channel(channels, chan, catch=False) - elif chan in channels: - try: - val[k] = read_channel(channels, chan, catch=False) - except EmptyChannelError: - continue - else: - val[k] = managed[k]() - except EmptyChannelError: - return + val: dict[str, Any] = {} + for k, chan in proc.channels.items(): + if chan in channels: + if channels[chan].is_available(): + val[k] = channels[chan].get() + else: + val[k] = managed[k]() elif isinstance(proc.channels, list): for chan in proc.channels: - try: - val = read_channel(channels, chan, catch=False) - break - except EmptyChannelError: - pass + if chan in channels: + if channels[chan].is_available(): + val = channels[chan].get() + break + else: + val[k] = managed[k]() else: - return + return MISSING else: raise RuntimeError( "Invalid channels type, expected list or dict, got {proc.channels}" @@ -963,7 +970,11 @@ def _proc_input( if for_execution and proc.mapper is not None: val = proc.mapper(val) - yield val + # Cache the input value + if input_cache is not None: + input_cache[proc.input_cache_key] = val + + return val def _uuid5_str(namespace: bytes, *parts: str) -> str: diff --git a/libs/langgraph/langgraph/pregel/read.py b/libs/langgraph/langgraph/pregel/read.py index 05d0c6b60..278ce1623 100644 --- a/libs/langgraph/langgraph/pregel/read.py +++ b/libs/langgraph/langgraph/pregel/read.py @@ -9,6 +9,7 @@ from typing import ( Mapping, Optional, Sequence, + TypeAlias, Union, ) @@ -30,6 +31,7 @@ from langgraph.utils.config import merge_configs from langgraph.utils.runnable import RunnableCallable, RunnableSeq READ_TYPE = Callable[[Union[str, Sequence[str]], bool], Union[Any, dict[str, Any]]] +INPUT_CACHE_KEY_TYPE: TypeAlias = tuple[Callable[..., Any], tuple[str, ...]] class ChannelRead(RunnableCallable): @@ -228,6 +230,17 @@ class PregelNode(Runnable): else: return self.bound + @cached_property + def input_cache_key(self) -> INPUT_CACHE_KEY_TYPE: + """Get a cache key for the input to the node. + This is used to avoid calculating the same input multiple times.""" + return ( + self.mapper, + tuple(f"{key}:{value}" for key, value in self.channels.items()) + if isinstance(self.channels, dict) + else tuple(self.channels), + ) + def join(self, channels: Sequence[str]) -> PregelNode: assert isinstance(channels, list) or isinstance( channels, tuple