From e7bc74e9186ca4d74bdcb14c9bf426f1d1acb723 Mon Sep 17 00:00:00 2001 From: vbarda Date: Wed, 21 Aug 2024 18:06:18 -0400 Subject: [PATCH] remove refactors --- libs/langgraph/langgraph/pregel/__init__.py | 163 +++++++++----------- 1 file changed, 70 insertions(+), 93 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index cb8cb29fb..5c54d5b7e 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -258,73 +258,6 @@ def _assemble_state_snapshot_hierarchy( return state_snapshot -def _prepare_state_snapshot( - saved: CheckpointTuple, - graph: Pregel, -) -> StateSnapshot: - with ChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in graph.channels.items() - }, - saved.checkpoint, - saved.config, - ) as channels, ManagedValuesManager( - graph.managed_values_dict, ensure_config(saved.config) - ) as managed: - next_tasks = prepare_next_tasks( - saved.checkpoint, - graph.nodes, - channels, - managed, - saved.config, - saved.metadata.get("step", -1) + 1, - for_execution=False, - ) - return StateSnapshot( - values=read_channels(channels, graph.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, - tasks=tasks_w_writes(next_tasks, saved.pending_writes), - ) - - -async def _prepare_state_snapshot_async( - saved: CheckpointTuple, graph: Pregel -) -> StateSnapshot: - async with AsyncChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in graph.channels.items() - }, - saved.checkpoint, - saved.config, - ) as channels, AsyncManagedValuesManager( - graph.managed_values_dict, ensure_config(saved.config) - ) as managed: - next_tasks = prepare_next_tasks( - saved.checkpoint, - graph.nodes, - channels, - managed, - saved.config, - saved.metadata.get("step", -1) + 1, - for_execution=False, - ) - return StateSnapshot( - values=read_channels(channels, graph.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, - tasks=tasks_w_writes(next_tasks, saved.pending_writes), - ) - - def _has_nested_interrupts( graph: Pregel, ) -> bool: @@ -517,20 +450,16 @@ class Pregel( if not self.checkpointer: raise ValueError("No checkpointer set") - checkpoint_tuple = self.checkpointer.get_tuple(config) - checkpoint_config = checkpoint_tuple.config if checkpoint_tuple else config + saved = self.checkpointer.get_tuple(config) + checkpoint_config = saved.config if saved else config checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "") checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id") checkpoint_ns_to_checkpoint_id: dict[str, str] = {} checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {} checkpoint_ns_to_graph: dict[str, Pregel] = {} - for checkpoint_tuple in self.checkpointer.list(config): - saved_checkpoint_ns = checkpoint_tuple.config["configurable"][ - "checkpoint_ns" - ] - saved_checkpoint_id = checkpoint_tuple.config["configurable"][ - "checkpoint_id" - ] + for saved in self.checkpointer.list(config): + saved_checkpoint_ns = saved.config["configurable"]["checkpoint_ns"] + saved_checkpoint_id = saved.config["configurable"]["checkpoint_id"] if checkpoint_id != saved_checkpoint_id: continue @@ -547,10 +476,36 @@ class Pregel( self, saved_checkpoint_ns ) - state_snapshot = _prepare_state_snapshot( - checkpoint_tuple, - checkpoint_ns_to_graph[saved_checkpoint_ns], - ) + graph = checkpoint_ns_to_graph[saved_checkpoint_ns] + with ChannelsManager( + { + k: LastValue(None) if isinstance(c, Context) else c + for k, c in graph.channels.items() + }, + saved.checkpoint, + saved.config, + ) as channels, ManagedValuesManager( + graph.managed_values_dict, ensure_config(saved.config) + ) as managed: + next_tasks = prepare_next_tasks( + saved.checkpoint, + graph.nodes, + channels, + managed, + saved.config, + saved.metadata.get("step", -1) + 1, + for_execution=False, + ) + state_snapshot = StateSnapshot( + read_channels(channels, graph.stream_channels_asis), + tuple(t.name for t in next_tasks), + saved.config, + saved.metadata, + saved.checkpoint["ts"], + saved.parent_config, + tasks_w_writes(next_tasks, saved.pending_writes), + ) + checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ saved_checkpoint_ns @@ -576,20 +531,16 @@ class Pregel( if not self.checkpointer: raise ValueError("No checkpointer set") - checkpoint_tuple = await self.checkpointer.aget_tuple(config) - checkpoint_config = checkpoint_tuple.config if checkpoint_tuple else config + saved = await self.checkpointer.aget_tuple(config) + checkpoint_config = saved.config if saved else config checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "") checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id") checkpoint_ns_to_checkpoint_id: dict[str, str] = {} checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {} checkpoint_ns_to_graph: dict[str, Pregel] = {} - async for checkpoint_tuple in self.checkpointer.alist(config): - saved_checkpoint_ns = checkpoint_tuple.config["configurable"][ - "checkpoint_ns" - ] - saved_checkpoint_id = checkpoint_tuple.config["configurable"][ - "checkpoint_id" - ] + async for saved in self.checkpointer.alist(config): + saved_checkpoint_ns = saved.config["configurable"]["checkpoint_ns"] + saved_checkpoint_id = saved.config["configurable"]["checkpoint_id"] if checkpoint_id != saved_checkpoint_id: continue @@ -606,10 +557,36 @@ class Pregel( self, saved_checkpoint_ns ) - state_snapshot = await _prepare_state_snapshot_async( - checkpoint_tuple, - checkpoint_ns_to_graph[saved_checkpoint_ns], - ) + graph = checkpoint_ns_to_graph[saved_checkpoint_ns] + async with AsyncChannelsManager( + { + k: LastValue(None) if isinstance(c, Context) else c + for k, c in graph.channels.items() + }, + saved.checkpoint, + saved.config, + ) as channels, AsyncManagedValuesManager( + graph.managed_values_dict, ensure_config(saved.config) + ) as managed: + next_tasks = prepare_next_tasks( + saved.checkpoint, + graph.nodes, + channels, + managed, + saved.config, + saved.metadata.get("step", -1) + 1, + for_execution=False, + ) + state_snapshot = StateSnapshot( + read_channels(channels, graph.stream_channels_asis), + tuple(t.name for t in next_tasks), + saved.config, + saved.metadata, + saved.checkpoint["ts"], + saved.parent_config, + tasks_w_writes(next_tasks, saved.pending_writes), + ) + checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot checkpoint_ns_to_checkpoint_id[ saved_checkpoint_ns