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 <noreply@anthropic.com>
This commit is contained in:
Sydney Runkle
2026-04-30 14:44:39 -04:00
co-authored by Claude Sonnet 4.6
parent 4856c89b30
commit f53ae61576
2 changed files with 68 additions and 12 deletions
+39 -4
View File
@@ -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__ != "<lambda>"
and other.operator.__name__ != "<lambda>"
@@ -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
@@ -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}"
)
# ---------------------------------------------------------------------------