diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 8fe9fed2e..40f0212cf 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -1443,9 +1443,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou step + 3, for_execution=True, store=self.store, - checkpointer=self.checkpointer - if isinstance(self.checkpointer, BaseCheckpointSaver) - else None, + checkpointer=checkpointer, manager=None, ) # apply null writes @@ -1455,10 +1453,10 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou if w[0] == NULL_TASK_ID ]: apply_writes( - saved.checkpoint, + checkpoint, channels, [PregelTaskWrites((), INPUT, null_writes, [])], - None, + checkpointer.get_next_version, self.trigger_to_nodes, ) # apply writes from tasks that already ran @@ -1473,19 +1471,22 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou checkpoint, channels, next_tasks.values(), - None, + checkpointer.get_next_version, self.trigger_to_nodes, ) # save checkpoint next_config = checkpointer.put( checkpoint_config, - create_checkpoint(checkpoint, None, step), + create_checkpoint(checkpoint, channels, step), { "source": "update", "step": step + 1, "parents": saved.metadata.get("parents", {}) if saved else {}, }, - {}, + get_new_channel_versions( + checkpoint_previous_versions, + checkpoint["channel_versions"], + ), ) return patch_checkpoint_map( next_config, saved.metadata if saved else None @@ -1645,11 +1646,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou step + 3, for_execution=True, store=self.store, - checkpointer=( - self.checkpointer - if isinstance(self.checkpointer, BaseCheckpointSaver) - else None - ), + checkpointer=checkpointer, manager=None, ) # apply null writes @@ -1657,10 +1654,10 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou w[1:] for w in saved.pending_writes or [] if w[0] == NULL_TASK_ID ]: apply_writes( - saved.checkpoint, + checkpoint, channels, [PregelTaskWrites((), INPUT, null_writes, [])], - None, + checkpointer.get_next_version, self.trigger_to_nodes, ) # apply writes @@ -1672,7 +1669,11 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou next_tasks[tid].writes.append((k, v)) if tasks := [t for t in next_tasks.values() if t.writes]: apply_writes( - checkpoint, channels, tasks, None, self.trigger_to_nodes + checkpoint, + channels, + tasks, + checkpointer.get_next_version, + self.trigger_to_nodes, ) valid_updates: list[tuple[str, dict[str, Any] | None, str | None]] = [] if len(updates) == 1: @@ -1901,9 +1902,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou step + 3, for_execution=True, store=self.store, - checkpointer=self.checkpointer - if isinstance(self.checkpointer, BaseCheckpointSaver) - else None, + checkpointer=checkpointer, manager=None, ) # apply null writes @@ -1913,10 +1912,10 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou if w[0] == NULL_TASK_ID ]: apply_writes( - saved.checkpoint, + checkpoint, channels, [PregelTaskWrites((), INPUT, null_writes, [])], - None, + checkpointer.get_next_version, self.trigger_to_nodes, ) # apply writes from tasks that already ran @@ -1931,19 +1930,21 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou checkpoint, channels, next_tasks.values(), - None, + checkpointer.get_next_version, self.trigger_to_nodes, ) # save checkpoint next_config = await checkpointer.aput( checkpoint_config, - create_checkpoint(checkpoint, None, step), + create_checkpoint(checkpoint, channels, step), { "source": "update", "step": step + 1, "parents": saved.metadata.get("parents", {}) if saved else {}, }, - {}, + get_new_channel_versions( + checkpoint_previous_versions, checkpoint["channel_versions"] + ), ) return patch_checkpoint_map( next_config, saved.metadata if saved else None @@ -2103,11 +2104,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou step + 3, for_execution=True, store=self.store, - checkpointer=( - self.checkpointer - if isinstance(self.checkpointer, BaseCheckpointSaver) - else None - ), + checkpointer=checkpointer, manager=None, ) # apply null writes @@ -2115,10 +2112,10 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou w[1:] for w in saved.pending_writes or [] if w[0] == NULL_TASK_ID ]: apply_writes( - saved.checkpoint, + checkpoint, channels, [PregelTaskWrites((), INPUT, null_writes, [])], - None, + checkpointer.get_next_version, self.trigger_to_nodes, ) for tid, k, v in saved.pending_writes: diff --git a/libs/langgraph/tests/test_large_cases.py b/libs/langgraph/tests/test_large_cases.py index cb24702ff..e3108c105 100644 --- a/libs/langgraph/tests/test_large_cases.py +++ b/libs/langgraph/tests/test_large_cases.py @@ -4454,15 +4454,9 @@ def test_partial_pending_checkpoint(sync_checkpointer: BaseCheckpointSaver) -> N # interrupt and unresolved tasks are cleared, finished tasks are kept assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value ⛰️", "market": "DE"}, - next=("tool_one",), - tasks=( - PregelTask( - id=AnyStr(), - name="tool_one", - path=("__pregel_push", 0, False), - ), - ), + values={"my_key": "value ⛰️ one", "market": "DE"}, + next=(), + tasks=(), config={ "configurable": { "thread_id": "1", diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 8c09a91a3..3fb05ed34 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -40,7 +40,7 @@ from langgraph.checkpoint.base import ( CheckpointTuple, ) from langgraph.checkpoint.memory import InMemorySaver -from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, PUSH, START +from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, START from langgraph.errors import InvalidUpdateError, NodeInterrupt, ParentCommand from langgraph.func import entrypoint, task from langgraph.graph import END, StateGraph @@ -1023,18 +1023,9 @@ async def test_partial_pending_checkpoint( # interrupt and next tasks are cleared, finished tasks are kept tup_upd = await tool_two.checkpointer.aget_tuple(thread1) assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value ⛰️", "market": "DE"}, - next=("tool_one",), - tasks=( - PregelTask( - AnyStr(), - "tool_one", - (PUSH, 0, False), - error=None, - interrupts=(), - state=None, - ), - ), + values={"my_key": "value ⛰️ one", "market": "DE"}, + next=(), + tasks=(), config=tup_upd.config, created_at=tup_upd.checkpoint["ts"], metadata={