mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-06 01:37:49 +02:00
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:
co-authored by
Claude Sonnet 4.6
parent
4856c89b30
commit
f53ae61576
@@ -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}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user