From f53ae61576d7bf78210c67339f14cf74e700a474 Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Fri, 17 Apr 2026 16:30:17 -0400 Subject: [PATCH] feat(channels): add rehydrate_every to DiffChannel for bounded chain traversal Periodic full-snapshot checkpoints cap chain depth, trading a small amount of extra storage for bounded reconstruction time. Co-Authored-By: Claude Sonnet 4.6 --- libs/langgraph/langgraph/channels/diff.py | 43 +++++++++++++++++-- .../tests/test_diff_channel_benchmark.py | 37 ++++++++++++---- 2 files changed, 68 insertions(+), 12 deletions(-) diff --git a/libs/langgraph/langgraph/channels/diff.py b/libs/langgraph/langgraph/channels/diff.py index 09569bb33..09caeb3df 100644 --- a/libs/langgraph/langgraph/channels/diff.py +++ b/libs/langgraph/langgraph/channels/diff.py @@ -31,12 +31,22 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): messages: Annotated[list[AnyMessage], DiffChannel(add_messages)] """ - __slots__ = ("value", "operator", "_pending", "_base_version", "_overwritten") + __slots__ = ( + "value", + "operator", + "rehydrate_every", + "_pending", + "_base_version", + "_overwritten", + "_steps_since_rehydrate", + ) def __init__( self, operator: Callable[[list[Value], Any], list[Value]], typ: type = list, + *, + rehydrate_every: int | None = None, ) -> None: typ = _strip_extras(typ) if typ in ( @@ -46,6 +56,7 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): typ = list super().__init__(typ) self.operator = operator + self.rehydrate_every = rehydrate_every try: self.value: list[Value] = typ() except Exception: @@ -53,10 +64,13 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): self._pending: list[Any] = [] self._base_version: str | None = None self._overwritten: bool = False + self._steps_since_rehydrate: int = 0 def __eq__(self, other: object) -> bool: if not isinstance(other, DiffChannel): return False + if self.rehydrate_every != other.rehydrate_every: + return False if ( self.operator.__name__ != "" and other.operator.__name__ != "" @@ -73,16 +87,17 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): return self.typ | list[self.typ] # type: ignore[name-defined] def copy(self) -> Self: - new = DiffChannel(self.operator, self.typ) + new = DiffChannel(self.operator, self.typ, rehydrate_every=self.rehydrate_every) new.key = self.key new.value = self.value[:] new._pending = self._pending[:] new._base_version = self._base_version new._overwritten = self._overwritten + new._steps_since_rehydrate = self._steps_since_rehydrate return new def from_checkpoint(self, checkpoint: Any) -> Self: - new = DiffChannel(self.operator, self.typ) + new = DiffChannel(self.operator, self.typ, rehydrate_every=self.rehydrate_every) new.key = self.key if checkpoint is MISSING: new.value = [] @@ -92,6 +107,9 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): for write in step_writes: accumulated = new.operator(accumulated, write) new.value = accumulated + # Seed the counter from actual chain depth so rehydration fires at + # the right time regardless of how many prior invocations there were. + new._steps_since_rehydrate = len(checkpoint.deltas) elif isinstance(checkpoint, DiffDelta): raise ValueError( "DiffChannel received a raw DiffDelta from the checkpoint saver. " @@ -144,7 +162,15 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): def is_available(self) -> bool: return self.value is not MISSING - def checkpoint(self) -> DiffDelta: + def checkpoint(self) -> Any: + if ( + self.rehydrate_every is not None + and self._steps_since_rehydrate >= self.rehydrate_every + ): + # Emit a full snapshot to cap chain depth at rehydrate_every. + # The saver stores this as a plain (non-diff) blob, so future + # deltas will chain back to it and traversal depth resets to 1. + return list(self.value) return DiffDelta( delta=self._pending[:], prev_version=None if self._overwritten else self._base_version, @@ -152,6 +178,15 @@ class DiffChannel(Generic[Value], BaseChannel[list[Value], Value, DiffDelta]): def after_checkpoint(self, version: Any) -> None: if version != self._base_version: + if self._base_version is None: + # First call after from_checkpoint — anchor the base version + # without counting a step (the counter was seeded by from_checkpoint). + pass + elif self.rehydrate_every is not None: + if self._steps_since_rehydrate >= self.rehydrate_every: + self._steps_since_rehydrate = 0 + else: + self._steps_since_rehydrate += 1 self._base_version = version self._pending = [] self._overwritten = False diff --git a/libs/langgraph/tests/test_diff_channel_benchmark.py b/libs/langgraph/tests/test_diff_channel_benchmark.py index 6de7c7fed..c5a88a40e 100644 --- a/libs/langgraph/tests/test_diff_channel_benchmark.py +++ b/libs/langgraph/tests/test_diff_channel_benchmark.py @@ -17,6 +17,8 @@ from langgraph.channels.diff import DiffChannel from langgraph.graph import END, StateGraph from langgraph.graph.message import add_messages +REHYDRATE_EVERY = 50 + # --------------------------------------------------------------------------- # State definitions @@ -31,6 +33,10 @@ class DiffState(TypedDict): messages: Annotated[list, DiffChannel(add_messages)] +class DiffRehydrateState(TypedDict): + messages: Annotated[list, DiffChannel(add_messages, rehydrate_every=REHYDRATE_EVERY)] + + # --------------------------------------------------------------------------- # Graph factory # --------------------------------------------------------------------------- @@ -91,23 +97,33 @@ TURN_COUNTS = [10, 50, 100, 200, 500] def run_benchmark() -> None: print() print("DiffChannel vs BinaryOperatorAggregate — checkpoint storage & time benchmark") - print("=" * 76) - header = f"{'turns':>6} {'binary_bytes':>14} {'diff_bytes':>12} {'ratio':>8} {'binary_ms':>12} {'diff_ms':>10}" + w = 100 + print("=" * w) + header = ( + f"{'turns':>6} " + f"{'bin_bytes':>12} {'diff_bytes':>12} {'rehy_bytes':>12} {'bytes_ratio':>12} " + f"{'bin_ms':>9} {'diff_ms':>9} {'rehy_ms':>9} {'time_ratio':>12}" + ) print(header) - print("-" * 76) + print("-" * w) for turns in TURN_COUNTS: b_time, b_bytes = _run_turns(turns, BinaryState) d_time, d_bytes = _run_turns(turns, DiffState) - ratio = b_bytes / d_bytes if d_bytes else float("inf") + r_time, r_bytes = _run_turns(turns, DiffRehydrateState) + bytes_ratio = b_bytes / d_bytes if d_bytes else float("inf") + time_ratio = d_time / b_time if b_time else float("inf") print( - f"{turns:>6} {b_bytes:>14,} {d_bytes:>12,} {ratio:>7.1f}x " - f"{b_time * 1000:>11.1f}ms {d_time * 1000:>9.1f}ms" + f"{turns:>6} " + f"{b_bytes:>12,} {d_bytes:>12,} {r_bytes:>12,} {bytes_ratio:>11.1f}x " + f"{b_time * 1000:>8.1f}ms {d_time * 1000:>8.1f}ms {r_time * 1000:>8.1f}ms {time_ratio:>11.1f}x" ) - print("=" * 76) + print("=" * w) print() - print("ratio = binary_bytes / diff_bytes (higher = DiffChannel saves more space)") + print(f"bytes_ratio = bin_bytes / diff_bytes (higher = more storage saved)") + print(f"time_ratio = diff_ms / bin_ms (higher = more overhead without rehydration)") + print(f"rehy = DiffChannel(rehydrate_every={REHYDRATE_EVERY}) — caps chain depth") print() @@ -125,10 +141,15 @@ def test_diff_channel_benchmark(capsys: Any) -> None: for turns in [100, 500]: _, b_bytes = _run_turns(turns, BinaryState) _, d_bytes = _run_turns(turns, DiffState) + _, r_bytes = _run_turns(turns, DiffRehydrateState) assert d_bytes < b_bytes, ( f"Expected DiffChannel to use less storage at {turns} turns, " f"got diff={d_bytes} binary={b_bytes}" ) + assert r_bytes < b_bytes, ( + f"Expected DiffChannel(rehydrate) to use less storage at {turns} turns, " + f"got rehydrate={r_bytes} binary={b_bytes}" + ) # ---------------------------------------------------------------------------