feat(pregel): wire DeltaChannel assembly into loop; pass checkpoint_id to after_checkpoint

This commit is contained in:
Sydney Runkle
2026-04-21 10:04:19 -04:00
parent 9fd6374302
commit c9913afef2
+25 -1
View File
@@ -92,6 +92,8 @@ from langgraph.pregel._algo import (
task_path_str,
)
from langgraph.pregel._checkpoint import (
_aassemble_delta_channels,
_assemble_delta_channels,
channels_from_checkpoint,
copy_checkpoint,
create_checkpoint,
@@ -883,7 +885,10 @@ class PregelLoop:
)
if do_checkpoint and self.channels:
for k, ch in self.channels.items():
ch.after_checkpoint(self.checkpoint["channel_versions"].get(k))
ch.after_checkpoint(
self.checkpoint["channel_versions"].get(k),
self.checkpoint.get("id"),
)
# sanitize TASK channel in the checkpoint before saving (durability=="exit")
if TASKS in self.checkpoint["channel_values"] and any(
isinstance(channel, UntrackedValue) for channel in self.channels.values()
@@ -1265,6 +1270,16 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
else []
)
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
# Assemble any DeltaChannel chains before constructing channel objects.
if self.checkpointer is not None:
assembled = _assemble_delta_channels(
self.checkpoint, self.checkpoint_config, self.checkpointer
)
if assembled:
self.checkpoint = {
**self.checkpoint,
"channel_values": {**self.checkpoint["channel_values"], **assembled},
}
self.channels, self.managed = channels_from_checkpoint(
self.specs, self.checkpoint
)
@@ -1469,6 +1484,15 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
self.submit = await self.stack.enter_async_context(
AsyncBackgroundExecutor(self.config)
)
if self.checkpointer is not None:
assembled = await _aassemble_delta_channels(
self.checkpoint, self.checkpoint_config, self.checkpointer
)
if assembled:
self.checkpoint = {
**self.checkpoint,
"channel_values": {**self.checkpoint["channel_values"], **assembled},
}
self.channels, self.managed = channels_from_checkpoint(
self.specs, self.checkpoint
)