From 72893d9abb9a813bd882d84f9f64cbe89227a938 Mon Sep 17 00:00:00 2001 From: vbarda Date: Thu, 22 Aug 2024 17:33:17 -0400 Subject: [PATCH] code review --- libs/langgraph/langgraph/pregel/__init__.py | 57 ++++----------------- libs/langgraph/langgraph/pregel/types.py | 2 +- libs/langgraph/langgraph/pregel/utils.py | 36 +++++++++++++ 3 files changed, 47 insertions(+), 48 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 5ee96dde5..7f15d01a4 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -88,7 +88,10 @@ from langgraph.pregel.types import ( StateSnapshot, StreamMode, ) -from langgraph.pregel.utils import get_new_channel_versions +from langgraph.pregel.utils import ( + assemble_state_snapshot_hierarchy, + get_new_channel_versions, +) from langgraph.pregel.validate import validate_graph, validate_keys from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore @@ -205,40 +208,6 @@ def _get_subgraph(graph: Pregel, checkpoint_ns: str) -> Pregel: return subgraph_node.bound -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 _has_nested_interrupts( graph: Pregel, ) -> bool: @@ -497,7 +466,7 @@ class Pregel( tasks=(), ) - state_snapshot = _assemble_state_snapshot_hierarchy( + state_snapshot = assemble_state_snapshot_hierarchy( checkpoint_ns, checkpoint_ns_to_state_snapshots ) return state_snapshot @@ -573,7 +542,7 @@ class Pregel( tasks=(), ) - state_snapshot = _assemble_state_snapshot_hierarchy( + state_snapshot = assemble_state_snapshot_hierarchy( checkpoint_ns, checkpoint_ns_to_state_snapshots ) return state_snapshot @@ -1113,7 +1082,8 @@ class Pregel( ) if not done: break # timed out - for fut, task in zip(done, [futures.pop(fut) for fut in done]): + for fut in done: + task = futures.pop(fut) if exc := _exception(fut): # save error to checkpointer if isinstance(exc, GraphInterrupt): @@ -1364,7 +1334,8 @@ class Pregel( if not done: break # timed out - for fut, task in zip(done, [futures.pop(fut) for fut in done]): + for fut in done: + task = futures.pop(fut) if exc := _exception(fut): # save error to checkpointer if isinstance(exc, GraphInterrupt): @@ -1581,11 +1552,3 @@ def _panic_or_proceed( inflight.pop().cancel() # raise timeout error raise timeout_exc_cls(f"Timed out at step {step}") - - -def _with_mode(mode: StreamMode, on: bool, iter: Iterator[Any]) -> Iterator[Any]: - if on: - for chunk in iter: - yield (mode, chunk) - else: - yield from iter diff --git a/libs/langgraph/langgraph/pregel/types.py b/libs/langgraph/langgraph/pregel/types.py index a8d66cb20..d4881452b 100644 --- a/libs/langgraph/langgraph/pregel/types.py +++ b/libs/langgraph/langgraph/pregel/types.py @@ -93,7 +93,7 @@ class StateSnapshot(NamedTuple): tasks: tuple[PregelTask, ...] """Tasks to execute in this step. If already attempted, may contain an error.""" subgraph_state_snapshots: Optional[dict[str, "StateSnapshot"]] = None - """State snapshots of subgraphs represented as a mapping from thread ID suffix to snapshot.""" + """State snapshots of subgraphs represented as a mapping from checkpoint namespace (`checkpoint_ns`) to snapshot.""" All = Literal["*"] diff --git a/libs/langgraph/langgraph/pregel/utils.py b/libs/langgraph/langgraph/pregel/utils.py index d3d0d989f..9c76e5f7e 100644 --- a/libs/langgraph/langgraph/pregel/utils.py +++ b/libs/langgraph/langgraph/pregel/utils.py @@ -1,4 +1,6 @@ from langgraph.checkpoint.base import ChannelVersions +from langgraph.constants import CHECKPOINT_NAMESPACE_SEPARATOR +from langgraph.pregel.types import StateSnapshot def get_new_channel_versions( @@ -17,3 +19,37 @@ def get_new_channel_versions( new_versions = current_versions return new_versions + + +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