diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index 89a2effa7..bc74619ff 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -86,7 +86,7 @@ class Send: self.id = id or str(uuid4()) def __hash__(self) -> int: - return hash((self.node, self.arg)) + return hash((self.node, self.arg, self.id)) def __repr__(self) -> str: return f"Send(node={self.node!r}, arg={self.arg!r}, id={self.id!r})" diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 81f08d98f..ffa53da85 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -195,6 +195,145 @@ class Channel: ) +def _get_nodes_and_channels( + nodes: Mapping[str, PregelNode], + channels: Mapping[str, BaseChannel], + checkpoint_ns: str, +) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]: + if checkpoint_ns == "": + return nodes, channels + + path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR) + for subgraph_node_name in path: + # if we have this separator it means we have a node that was triggered by Send + if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name: + name_parts = subgraph_node_name.split(SEND_CHECKPOINT_NAMESPACE_SEPARATOR) + if len(name_parts) != 2: + raise ValueError(f"Malformed node name '{subgraph_node_name}'") + + subgraph_node_name = name_parts[0] + if subgraph_node_name not in nodes: + raise ValueError(f"Couldn't find node '{subgraph_node_name}'.") + + subgraph_node = nodes[subgraph_node_name].get_node() + + if not isinstance(subgraph_node, RunnableSequence): + break + + first_step = subgraph_node.steps[0] + if isinstance(first_step, Pregel): + nodes = first_step.nodes + channels = first_step.channels + + return nodes, channels + + +def _assemble_state_snapshot_hierarchy( + root_checkpoint_ns: str, + checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot], +) -> StateSnapshot: + checkpoint_ns_list_to_visit = sorted( + checkpoint_ns_to_state_snapshots.keys(), + key=lambda x: len(x.split(CHECKPOINT_NAMESPACE_SEPARATOR)), + ) + while checkpoint_ns_list_to_visit: + checkpoint_ns = checkpoint_ns_list_to_visit.pop() + state_snapshot = checkpoint_ns_to_state_snapshots[checkpoint_ns] + *path, subgraph_node = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR) + parent_checkpoint_ns = CHECKPOINT_NAMESPACE_SEPARATOR.join(path) + if subgraph_node and ( + parent_state_snapshot := checkpoint_ns_to_state_snapshots.get( + parent_checkpoint_ns + ) + ): + parent_subgraph_snapshots = { + **(parent_state_snapshot.subgraph_state_snapshots or {}), + subgraph_node: state_snapshot, + } + checkpoint_ns_to_state_snapshots[ + parent_checkpoint_ns + ] = checkpoint_ns_to_state_snapshots[parent_checkpoint_ns]._replace( + subgraph_state_snapshots=parent_subgraph_snapshots + ) + + state_snapshot = checkpoint_ns_to_state_snapshots.pop(root_checkpoint_ns, None) + if state_snapshot is None: + raise ValueError(f"Missing checkpoint for checkpoint NS '{root_checkpoint_ns}'") + return state_snapshot + + +def _prepare_state_snapshot( + saved: CheckpointTuple, + nodes: Mapping[str, PregelNode], + channels: Mapping[str, BaseChannel], + managed_values_dict: dict[str, ManagedValueSpec], + select_channels: str | list[str], +) -> StateSnapshot: + with ChannelsManager( + { + k: LastValue(None) if isinstance(c, Context) else c + for k, c in channels.items() + }, + saved.checkpoint, + saved.config, + ) as channels, ManagedValuesManager( + managed_values_dict, ensure_config(saved.config) + ) as managed: + next_tasks = prepare_next_tasks( + saved.checkpoint, + nodes, + channels, + managed, + saved.config, + -1, + for_execution=False, + ) + return StateSnapshot( + values=read_channels(channels, select_channels), + next=tuple(t.name for t in next_tasks), + config=saved.config, + metadata=saved.metadata, + created_at=saved.checkpoint["ts"], + parent_config=saved.parent_config, + ) + + +async def _prepare_state_snapshot_async( + saved: CheckpointTuple, + nodes: Mapping[str, PregelNode], + channels: Mapping[str, BaseChannel], + managed_values_dict: dict[str, ManagedValueSpec], + select_channels: str | list[str], +) -> StateSnapshot: + async with AsyncChannelsManager( + { + k: LastValue(None) if isinstance(c, Context) else c + for k, c in channels.items() + }, + saved.checkpoint, + saved.config, + ) as channels, AsyncManagedValuesManager( + managed_values_dict, ensure_config(saved.config) + ) as managed: + next_tasks = prepare_next_tasks( + saved.checkpoint, + nodes, + channels, + managed, + saved.config, + -1, + for_execution=False, + ) + return StateSnapshot( + values=read_channels(channels, select_channels), + next=tuple(t.name for t in next_tasks), + config=saved.config, + metadata=saved.metadata, + created_at=saved.checkpoint["ts"], + parent_config=saved.parent_config, + ) + + class Pregel( RunnableSerializable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]] ): @@ -361,144 +500,6 @@ class Pregel( if is_managed_value(v) } - def _get_nodes_and_channels( - self, checkpoint_ns: str - ) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]: - if checkpoint_ns == "": - return self.nodes, self.channels - - path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR) - nodes = self.nodes - channels = self.channels - for subgraph_node_name in path: - # if we have this separator it means we have a node that was triggered by Send - if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name: - name_parts = subgraph_node_name.split( - SEND_CHECKPOINT_NAMESPACE_SEPARATOR - ) - if len(name_parts) != 2: - raise ValueError(f"Malformed node name '{subgraph_node_name}'") - - subgraph_node_name = name_parts[0] - if subgraph_node_name not in nodes: - raise ValueError(f"Couldn't find node '{subgraph_node_name}'.") - - subgraph_node = nodes[subgraph_node_name].get_node() - - if not isinstance(subgraph_node, RunnableSequence): - break - - first_step = subgraph_node.steps[0] - if isinstance(first_step, Pregel): - nodes = first_step.nodes - channels = first_step.channels - - return nodes, channels - - def _prepare_state_snapshot( - self, - saved: CheckpointTuple, - nodes: Mapping[str, PregelNode], - channels: Mapping[str, BaseChannel], - ) -> StateSnapshot: - with ChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in channels.items() - }, - saved.checkpoint, - saved.config, - ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(saved.config) - ) as managed: - next_tasks = prepare_next_tasks( - saved.checkpoint, - nodes, - channels, - managed, - saved.config, - -1, - for_execution=False, - ) - return StateSnapshot( - values=read_channels(channels, self.stream_channels_asis), - next=tuple(t.name for t in next_tasks), - config=saved.config, - metadata=saved.metadata, - created_at=saved.checkpoint["ts"], - parent_config=saved.parent_config, - ) - - async def _prepare_state_snapshot_async( - self, - saved: CheckpointTuple, - nodes: Mapping[str, PregelNode], - channels: Mapping[str, BaseChannel], - ) -> StateSnapshot: - async with AsyncChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in channels.items() - }, - saved.checkpoint, - saved.config, - ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(saved.config) - ) as managed: - next_tasks = prepare_next_tasks( - saved.checkpoint, - nodes, - channels, - managed, - saved.config, - -1, - for_execution=False, - ) - return StateSnapshot( - values=read_channels(channels, self.stream_channels_asis), - next=tuple(t.name for t in next_tasks), - config=saved.config, - metadata=saved.metadata, - created_at=saved.checkpoint["ts"], - parent_config=saved.parent_config, - ) - - @staticmethod - def _assemble_state_snapshot_hierarchy( - root_checkpoint_ns: str, - checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot], - ) -> StateSnapshot: - checkpoint_ns_list_to_visit = sorted( - checkpoint_ns_to_state_snapshots.keys(), - key=lambda x: len(x.split(CHECKPOINT_NAMESPACE_SEPARATOR)), - ) - while checkpoint_ns_list_to_visit: - checkpoint_ns = checkpoint_ns_list_to_visit.pop() - state_snapshot = checkpoint_ns_to_state_snapshots[checkpoint_ns] - *path, subgraph_node = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR) - parent_checkpoint_ns = CHECKPOINT_NAMESPACE_SEPARATOR.join(path) - if subgraph_node and ( - parent_state_snapshot := checkpoint_ns_to_state_snapshots.get( - parent_checkpoint_ns - ) - ): - parent_subgraph_snapshots = { - **(parent_state_snapshot.subgraph_state_snapshots or {}), - subgraph_node: state_snapshot, - } - checkpoint_ns_to_state_snapshots[ - parent_checkpoint_ns - ] = checkpoint_ns_to_state_snapshots[parent_checkpoint_ns]._replace( - subgraph_state_snapshots=parent_subgraph_snapshots - ) - - state_snapshot = checkpoint_ns_to_state_snapshots.pop(root_checkpoint_ns, None) - if state_snapshot is None: - raise ValueError( - f"Missing checkpoint for checkpoint NS '{root_checkpoint_ns}'" - ) - return state_snapshot - def get_state( self, config: RunnableConfig, *, include_subgraph_state: bool = False ) -> StateSnapshot: @@ -541,13 +542,19 @@ class Pregel( if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels: checkpoint_ns_to_nodes_and_channels[ saved_checkpoint_ns - ] = self._get_nodes_and_channels(saved_checkpoint_ns) + ] = _get_nodes_and_channels( + self.nodes, self.channels, saved_checkpoint_ns + ) nodes, channels = checkpoint_ns_to_nodes_and_channels[ saved_checkpoint_ns ] - state_snapshot = self._prepare_state_snapshot( - checkpoint_tuple, nodes, channels + state_snapshot = _prepare_state_snapshot( + checkpoint_tuple, + nodes, + channels, + self.managed_values_dict, + self.stream_channels_asis, ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ @@ -559,7 +566,7 @@ class Pregel( values={}, next=(), config=config, metadata=None, created_at=None ) - state_snapshot = self._assemble_state_snapshot_hierarchy( + state_snapshot = _assemble_state_snapshot_hierarchy( checkpoint_ns, checkpoint_ns_to_state_snapshots ) return state_snapshot @@ -611,13 +618,19 @@ class Pregel( if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels: checkpoint_ns_to_nodes_and_channels[ saved_checkpoint_ns - ] = self._get_nodes_and_channels(saved_checkpoint_ns) + ] = _get_nodes_and_channels( + self.nodes, self.channels, saved_checkpoint_ns + ) nodes, channels = checkpoint_ns_to_nodes_and_channels[ saved_checkpoint_ns ] - state_snapshot = await self._prepare_state_snapshot_async( - checkpoint_tuple, nodes, channels + state_snapshot = await _prepare_state_snapshot_async( + checkpoint_tuple, + nodes, + channels, + self.managed_values_dict, + self.stream_channels_asis, ) checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ @@ -629,7 +642,7 @@ class Pregel( values={}, next=(), config=config, metadata=None, created_at=None ) - state_snapshot = self._assemble_state_snapshot_hierarchy( + state_snapshot = _assemble_state_snapshot_hierarchy( checkpoint_ns, checkpoint_ns_to_state_snapshots ) return state_snapshot @@ -669,10 +682,18 @@ class Pregel( ) yield state_snapshot else: - nodes, channels = self._get_nodes_and_channels( - checkpoint_tuple.config["configurable"]["checkpoint_ns"] + nodes, channels = _get_nodes_and_channels( + self.nodes, + self.channels, + checkpoint_tuple.config["configurable"]["checkpoint_ns"], + ) + yield _prepare_state_snapshot( + checkpoint_tuple, + nodes, + channels, + self.managed_values_dict, + self.stream_channels_asis, ) - yield self._prepare_state_snapshot(checkpoint_tuple, nodes, channels) async def aget_state_history( self, @@ -709,11 +730,17 @@ class Pregel( ) yield state_snapshot else: - nodes, channels = self._get_nodes_and_channels( - checkpoint_tuple.config["configurable"]["checkpoint_ns"] + nodes, channels = _get_nodes_and_channels( + self.nodes, + self.channels, + checkpoint_tuple.config["configurable"]["checkpoint_ns"], ) - yield await self._prepare_state_snapshot_async( - checkpoint_tuple, nodes, channels + yield await _prepare_state_snapshot_async( + checkpoint_tuple, + nodes, + channels, + self.managed_values_dict, + self.stream_channels_asis, ) def update_state(