Compare commits

...
Author SHA1 Message Date
Sydney Runkle bfec964248 refactor(delta-channel): infer typ from Annotated outer type instead of constructor arg
Remove the positional `typ` parameter from `DeltaChannel.__init__`. The type is
now injected automatically from the `Annotated` outer type in `_is_field_channel`
(matching how `BinaryOperatorAggregate` receives its type). `copy()` and
`from_checkpoint()` propagate `self.typ` explicitly. Test helpers updated to
use `_get_channel` with the proper `Annotated` path.
2026-04-22 10:46:17 -04:00
Sydney Runkle 57be304f55 chore(delta-channel): remove supports_delta_channels flag
Rely on the runtime raise in DeltaChannel.from_checkpoint() instead of
a compile-time boolean flag. Savers that assemble DeltaChainValue inside
_load_blobs work transparently; savers that don't will pass through a raw
DeltaValue and hit a clear ValueError on first reload.

Removes: BaseCheckpointSaver.supports_delta_channels, the attribute on
InMemorySaver / PostgresSaver / AsyncPostgresSaver, the compile-time
UserWarning in StateGraph.compile(), and the associated test.
2026-04-22 09:40:00 -04:00
Sydney RunkleandClaude Sonnet 4.6 e7fd2fc327 chore: remove docs/ from PR
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-22 09:07:16 -04:00
ccurmeandSydney Runkle 024cf29274 fix(prebuilt): handle injected NotRequired keys (#7392)
Resolves https://github.com/langchain-ai/langchain/issues/35585

This would previously raise KeyError:
```python
from typing import Annotated

from langchain_core.tools import tool
from langchain.agents import create_agent
from typing_extensions import NotRequired
from langgraph.prebuilt import InjectedState
from langchain.agents import AgentState


class CustomAgentState(AgentState):
    city: NotRequired[str]


@tool
def get_weather(city: Annotated[str | None, InjectedState("city")] = None) -> str:
    """Get weather for a given city."""
    if city is None:
        city = "Boston"
    return f"It's always sunny in {city}!"


agent = create_agent(
    model="claude-sonnet-4-6",
    tools=[get_weather],
    system_prompt="You are a helpful assistant",
    state_schema=CustomAgentState,
)

input_message = {
    "role": "user",
    "content": "What's the weather?",
}

result = agent.invoke({"messages": [input_message]})
for m in result["messages"]:
    m.pretty_print()
```

---------

Co-authored-by: Sydney Runkle <sydneymarierunkle@gmail.com>
2026-04-22 09:06:45 -04:00
Sydney Runkle 792f779c18 lint 2026-04-22 08:54:58 -04:00
Sydney Runkle 1099cece53 latest 2026-04-22 08:48:53 -04:00
Sydney Runkle 43ecd9dc2d fix(delta-channel): support non-list reducers (dict) and fix MISSING handling
Use typ() instead of [] throughout DeltaChannel so reducers over dict
(and other non-list types) work correctly. fromCheckpoint(MISSING) now
leaves value as typ() from __init__ instead of overwriting with MISSING.
copy() uses value.copy() to handle dicts. update() initialises base from
typ() when value is MISSING. Add four tests covering the deepagents-style
dict-merge / file-deletion reducer pattern.
2026-04-21 13:04:40 -04:00
Sydney RunkleandClaude Sonnet 4.6 c459079e52 chore(delta): rename _steps_since_rehydrate → _steps_since_snapshot; add audit tests
- Rename `_steps_since_rehydrate` → `_steps_since_snapshot` in DeltaChannel
  for clarity (counts steps since the last snapshot, not since rehydration)
- Pre-seed cycle-detection `visited` set with current checkpoint ID in both
  sync and async `_assemble_delta_channels` to prevent self-referential chains
- Add 4 new unit tests:
  - `test_delta_channel_snapshot_every_emits_plain_list`: verifies counter
    semantics and snapshot/delta transitions
  - `test_delta_channel_snapshot_every_end_to_end`: graph-level smoke test
  - `test_delta_channel_assembly_fast_path_returns_delta_value`: exercises
    chain traversal via get_channel_blob returning DeltaValue then plain list
  - `test_delta_channel_assembly_broken_chain_logs_warning`: partial chain
    when get_tuple returns None

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 11:03:58 -04:00
Sydney RunkleandClaude Sonnet 4.6 ffda8b5472 chore: rename serde type tag "diff" → "delta" for DeltaValue
Consistent with channel/type naming (DeltaChannel, DeltaValue).

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 10:58:09 -04:00
Sydney RunkleandClaude Sonnet 4.6 374eebcd65 chore: apply format/lint fixes across checkpoint, checkpoint-postgres, prebuilt
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 10:50:27 -04:00
Sydney RunkleandClaude Sonnet 4.6 350182ed18 fix: register DeltaValue in SAFE_MSGPACK_TYPES; rename _is_diff_delta; cross-saver benchmark
- Add DeltaValue to SAFE_MSGPACK_TYPES so SQLite and other msgpack-based
  savers don't emit "Deserializing unregistered type" warnings.
- Rename _is_diff_delta → _is_delta_value (leftover from DiffChannel rename).
- Parametrize benchmark by checkpointer: runs InMemory (fast-path) and
  SQLite (get_tuple fallback) in the same table, sharing the _run_turns helper.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 10:48:20 -04:00
Sydney RunkleandClaude Sonnet 4.6 4c2ce5c8a9 fix(delta-channel): fix chain assembly and get_state paths
- Fix InMemorySaver.get_channel_blob: use correct storage[thread_id][ns]
  nesting and deserialize the checkpoint before extracting channel_versions.
- Pass checkpoint_id to after_checkpoint() in channels_from_checkpoint so
  DeltaChannel seeds _last_checkpoint_id correctly on load; without this
  every turn broke the chain at its boundary.
- Wire _assemble_delta_channels into _prepare_state_snapshot and
  _aprepare_state_snapshot (get_state / get_state_history paths) and into
  perform_superstep / aperform_superstep (update_state paths) — previously
  only the loop __enter__ path did assembly.
- Fix test_get_channel_blob to use the correct storage structure.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 10:40:00 -04:00
Sydney Runkle bae7486565 test(channels): replace unsupported-saver raise test with fallback assembly test 2026-04-21 10:06:51 -04:00
Sydney Runkle b082584500 feat(postgres): remove _load_diff_chains; add get_channel_blob / aget_channel_blob 2026-04-21 10:06:13 -04:00
Sydney Runkle 503071c2aa feat(memory): implement get_channel_blob; remove diff handling from _load_blobs 2026-04-21 10:05:08 -04:00
Sydney Runkle c9913afef2 feat(pregel): wire DeltaChannel assembly into loop; pass checkpoint_id to after_checkpoint 2026-04-21 10:04:19 -04:00
Sydney Runkle 9fd6374302 feat(pregel): add _assemble_delta_channels helpers for universal DeltaChannel support 2026-04-21 09:37:49 -04:00
Sydney RunkleandClaude Sonnet 4.6 e6c065739f feat(channels): DeltaChannel tracks checkpoint_id; emits prev_checkpoint_id in DeltaValue
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 09:34:44 -04:00
Sydney RunkleandClaude Sonnet 4.6 e18f8fff2b feat(serde): diff type encodes prev_checkpoint_id; loads_typed returns DeltaValue
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-21 09:29:58 -04:00
Sydney Runkle 65438610e8 docs(checkpoint): expand aget_channel_blob docstring for parity 2026-04-21 09:28:21 -04:00
Sydney Runkle 4ebebd686a chore: add .worktrees/ to .gitignore 2026-04-21 09:27:47 -04:00
Sydney Runkle fca3f6d919 feat(checkpoint): DeltaValue uses prev_checkpoint_id; add get_channel_blob stubs 2026-04-21 09:27:22 -04:00
Sydney Runkle 599afd7585 chore: rename DiffChannel/DiffDelta/DiffChainValue to Delta* across libs
Renames the diff-channel types to DeltaChannel, DeltaValue, and DeltaChainValue
for consistency with the settled naming convention.
2026-04-21 07:56:43 -04:00
Sydney Runkle c0e6062bfb more tests 2026-04-20 12:53:54 -04:00
Sydney RunkleandClaude Sonnet 4.6 056d3143ff chore: format/lint fixes for rehydrate_every benchmark
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 16:31:26 -04:00
Sydney RunkleandClaude Sonnet 4.6 1da43d412b 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>
2026-04-17 16:30:17 -04:00
Sydney RunkleandClaude Sonnet 4.6 df56b7cdf6 test(channels): add DiffChannel vs BinaryOperatorAggregate storage/time benchmark
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 16:21:28 -04:00
Sydney RunkleandClaude Sonnet 4.6 566a3150b2 chore: format and lint fixes for DiffChannel implementation
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 16:20:21 -04:00
Sydney RunkleandClaude Sonnet 4.6 1bb1811fbd fix(checkpoint/postgres): pass cursor to avoid deadlock in diff chain traversal
Fixes a critical deadlock that occurs when _load_diff_chains calls self._cursor()
from within _load_blobs while the outer _load_checkpoint_tuple already holds
self._cursor(). On bare (non-pool) connections, the threading.Lock is not
reentrant, causing a deadlock.

Solution: Pass the cursor as a parameter to _load_diff_chains and _load_blobs
instead of acquiring a new cursor within those methods. Updated _load_checkpoint_tuple
to acquire a cursor once at the top level and pass it through the call chain.

Changes:
- Updated _load_blobs signature to accept optional cur parameter
- Updated _load_diff_chains signature (base and implementations) to accept optional cur parameter
- Modified _load_checkpoint_tuple in PostgresSaver to acquire cursor and pass it
- Modified _load_checkpoint_tuple_async to acquire cursor only when diff_payloads exist
- Removed nested self._cursor() calls in _load_diff_chains and _load_diff_chains_async

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 16:05:58 -04:00
Sydney RunkleandClaude Sonnet 4.6 dd7f21e3ff feat(checkpoint/postgres): diff chain reconstruction in async saver
Add `_load_diff_chains_async` to `AsyncPostgresSaver` and override
`_load_checkpoint_tuple` to inline blob-parsing and diff-chain
resolution via async point-lookup traversal, mirroring the sync
`PostgresSaver._load_diff_chains` implementation. Add integration test
`test_diff_channel_chain_reconstruction` that skips gracefully when
`langgraph` core is not installed.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 15:56:19 -04:00
Sydney RunkleandClaude Sonnet 4.6 d4e1efa1f6 feat(checkpoint/postgres): diff chain reconstruction in _load_blobs (sync)
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 15:37:18 -04:00
Sydney RunkleandClaude Sonnet 4.6 dba1987c9b test(pregel): strengthen DiffChannel time-travel and reply assertions
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 15:01:50 -04:00
Sydney RunkleandClaude Sonnet 4.6 9fb0493ac5 feat(pregel): call after_checkpoint hook when loading and saving channels
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:56:42 -04:00
Sydney RunkleandClaude Sonnet 4.6 1cb057e6bc fix(checkpoint/memory): warn on broken diff chain, guard against cycles
- Add logger.warning when a mid-chain blob is missing (fixes silent truncation bug)
- Add cycle guard to prevent infinite loops on corrupt blob stores
- Fix type annotation on diff_channels from dict[str, Any] to dict[str, str]

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:54:58 -04:00
Sydney RunkleandClaude Sonnet 4.6 d76127fbbf feat(checkpoint/memory): chain-traverse diff blobs in _load_blobs
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:52:34 -04:00
Sydney Runkle 94853fb14c fix(channels): align DiffChannel.is_available with BinaryOperatorAggregate 2026-04-17 14:51:03 -04:00
Sydney RunkleandClaude Sonnet 4.6 e6fab22f0c feat(channels): implement DiffChannel for incremental checkpoint storage
Adds DiffChannel, a new channel type that stores only per-step write
deltas in checkpoints and reconstructs the full list by replaying the
chain through the operator at load time.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:35:46 -04:00
Sydney Runkle ea644f413d feat(channels): add no-op after_checkpoint hook to BaseChannel 2026-04-17 14:30:04 -04:00
Sydney RunkleandClaude Sonnet 4.6 3f86b1485d fix(checkpoint/serde): use lazy isinstance check for DiffDelta
Replace duck-typing check with lazy import inside _is_diff_delta helper
function to avoid module-level circular dependency while using proper
isinstance semantics.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:24:41 -04:00
Sydney RunkleandClaude Sonnet 4.5 9e40dee07f feat(checkpoint/serde): serialize DiffDelta as 'diff' type tag
Add serde support for DiffDelta by implementing dump/load for the "diff" type tag.
This allows the checkpoint system to efficiently store delta objects by serializing
them as msgpack-encoded dicts with {"d": delta, "p": prev_version} structure.

The implementation uses runtime type checking to avoid circular imports and
leverages the existing msgpack ext hooks for proper deserialization of complex
types like LangChain messages.

Co-Authored-By: Claude Sonnet 4.5 <noreply@anthropic.com>
2026-04-17 14:21:40 -04:00
Sydney RunkleandClaude Sonnet 4.6 4d1f4086eb feat(checkpoint): add DiffDelta and DiffChainValue protocol types
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:18:09 -04:00
Sydney RunkleandClaude Sonnet 4.6 afcf6c03dd docs: add DiffChannel implementation plan
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:14:50 -04:00
Sydney RunkleandClaude Sonnet 4.6 4b303ceb39 docs: add DiffChannel incremental checkpoint storage design spec
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-04-17 14:00:41 -04:00
22 changed files with 1853 additions and 39 deletions
+1
View File
@@ -100,3 +100,4 @@ dmypy.json
.turbo
.editorconfig
.scratch
.worktrees/
@@ -32,7 +32,7 @@ Conn = _internal.Conn # For backward compatibility
class PostgresSaver(BasePostgresSaver):
"""Checkpointer that stores checkpoints in a Postgres database."""
lock: threading.Lock
lock: threading.RLock
def __init__(
self,
@@ -48,7 +48,7 @@ class PostgresSaver(BasePostgresSaver):
self.conn = conn
self.pipe = pipe
self.lock = threading.Lock()
self.lock = threading.RLock()
self.supports_pipeline = Capabilities().has_pipeline()
@classmethod
@@ -442,6 +442,14 @@ class PostgresSaver(BasePostgresSaver):
including its configuration, metadata, parent checkpoint (if any),
and pending writes.
"""
with self._cursor() as cur:
channel_values = self._load_blobs(
value["channel_values"],
thread_id=value["thread_id"],
checkpoint_ns=value["checkpoint_ns"],
checkpoint_id=value["checkpoint_id"],
cur=cur,
)
return CheckpointTuple(
{
"configurable": {
@@ -454,7 +462,7 @@ class PostgresSaver(BasePostgresSaver):
**value["checkpoint"],
"channel_values": {
**(value["checkpoint"].get("channel_values") or {}),
**self._load_blobs(value["channel_values"]),
**channel_values,
},
},
value["metadata"],
@@ -13,6 +13,8 @@ from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
DeltaChainValue,
DeltaValue,
get_checkpoint_id,
get_serializable_checkpoint_metadata,
)
@@ -391,6 +393,81 @@ class AsyncPostgresSaver(BasePostgresSaver):
async with conn.cursor(binary=True, row_factory=dict_row) as cur:
yield cur
async def _aload_delta_chain(
self,
thread_id: str,
checkpoint_ns: str,
checkpoint_id: str,
channel: str,
cur: Any,
) -> DeltaChainValue:
"""Fetch the full delta chain for a channel in one recursive CTE query (async)."""
await cur.execute(
"""
WITH RECURSIVE chain AS (
SELECT
c.checkpoint_id,
c.parent_checkpoint_id,
cb.version,
cb.type,
cb.blob
FROM checkpoints c
JOIN checkpoint_blobs cb
ON cb.thread_id = %s
AND cb.checkpoint_ns = %s
AND cb.channel = %s
AND cb.version = (c.checkpoint->'channel_versions'->>%s)::text
WHERE c.thread_id = %s
AND c.checkpoint_ns = %s
AND c.checkpoint_id = %s
UNION ALL
SELECT
c.checkpoint_id,
c.parent_checkpoint_id,
cb.version,
cb.type,
cb.blob
FROM chain prev
JOIN checkpoints c ON c.checkpoint_id = prev.parent_checkpoint_id
JOIN checkpoint_blobs cb
ON cb.thread_id = %s
AND cb.checkpoint_ns = %s
AND cb.channel = %s
AND cb.version = (c.checkpoint->'channel_versions'->>%s)::text
WHERE prev.parent_checkpoint_id IS NOT NULL
AND prev.type = 'delta'
)
SELECT DISTINCT ON (version) type, blob
FROM chain
ORDER BY version ASC
""",
(
thread_id,
checkpoint_ns,
channel,
channel,
thread_id,
checkpoint_ns,
checkpoint_id,
thread_id,
checkpoint_ns,
channel,
channel,
),
)
rows = await cur.fetchall()
base = None
deltas: list[list[Any]] = []
for row in rows:
blob = self.serde.loads_typed((row["type"], row["blob"]))
if isinstance(blob, DeltaValue):
deltas.append(blob.delta)
else:
base = blob
return DeltaChainValue(base=base, deltas=deltas)
async def _load_checkpoint_tuple(self, value: DictRow) -> CheckpointTuple:
"""
Convert a database row into a CheckpointTuple object.
@@ -403,11 +480,29 @@ class AsyncPostgresSaver(BasePostgresSaver):
including its configuration, metadata, parent checkpoint (if any),
and pending writes.
"""
thread_id = value["thread_id"]
checkpoint_ns = value["checkpoint_ns"]
checkpoint_id = value["checkpoint_id"]
blob_values = value["channel_values"]
channel_values: dict[str, Any] = {}
if blob_values:
channel_values = self._load_blobs(blob_values)
delta_channels = [
k.decode() for k, t, _ in blob_values if t.decode() == "delta"
]
if delta_channels:
async with self._cursor() as cur:
for channel in delta_channels:
channel_values[channel] = await self._aload_delta_chain(
thread_id, checkpoint_ns, checkpoint_id, channel, cur
)
return CheckpointTuple(
{
"configurable": {
"thread_id": value["thread_id"],
"checkpoint_ns": value["checkpoint_ns"],
"thread_id": thread_id,
"checkpoint_ns": checkpoint_ns,
"checkpoint_id": value["checkpoint_id"],
}
},
@@ -415,15 +510,15 @@ class AsyncPostgresSaver(BasePostgresSaver):
**value["checkpoint"],
"channel_values": {
**(value["checkpoint"].get("channel_values") or {}),
**self._load_blobs(value["channel_values"]),
**channel_values,
},
},
value["metadata"],
(
{
"configurable": {
"thread_id": value["thread_id"],
"checkpoint_ns": value["checkpoint_ns"],
"thread_id": thread_id,
"checkpoint_ns": checkpoint_ns,
"checkpoint_id": value["parent_checkpoint_id"],
}
}
@@ -11,6 +11,8 @@ from langgraph.checkpoint.base import (
WRITES_IDX_MAP,
BaseCheckpointSaver,
ChannelVersions,
DeltaChainValue,
DeltaValue,
get_checkpoint_id,
)
from langgraph.checkpoint.serde.types import TASKS
@@ -185,15 +187,106 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
)
def _load_blobs(
self, blob_values: list[tuple[bytes, bytes, bytes]]
self,
blob_values: list[tuple[bytes, bytes, bytes]],
*,
thread_id: str = "",
checkpoint_ns: str = "",
checkpoint_id: str = "",
cur: Any = None,
) -> dict[str, Any]:
if not blob_values:
return {}
return {
k.decode(): self.serde.loads_typed((t.decode(), v))
for k, t, v in blob_values
if t.decode() != "empty"
}
result: dict[str, Any] = {}
delta_channels: list[str] = []
for k, t, v in blob_values:
channel = k.decode()
type_tag = t.decode()
if type_tag == "delta":
delta_channels.append(channel)
elif type_tag != "empty":
result[channel] = self.serde.loads_typed((type_tag, v))
if delta_channels and cur is not None and checkpoint_id:
for channel in delta_channels:
result[channel] = self._load_delta_chain(
thread_id, checkpoint_ns, checkpoint_id, channel, cur
)
return result
def _load_delta_chain(
self,
thread_id: str,
checkpoint_ns: str,
checkpoint_id: str,
channel: str,
cur: Any,
) -> DeltaChainValue:
"""Fetch the full delta chain for a channel in one recursive CTE query."""
cur.execute(
"""
WITH RECURSIVE chain AS (
SELECT
c.checkpoint_id,
c.parent_checkpoint_id,
cb.version,
cb.type,
cb.blob
FROM checkpoints c
JOIN checkpoint_blobs cb
ON cb.thread_id = %s
AND cb.checkpoint_ns = %s
AND cb.channel = %s
AND cb.version = (c.checkpoint->'channel_versions'->>%s)::text
WHERE c.thread_id = %s
AND c.checkpoint_ns = %s
AND c.checkpoint_id = %s
UNION ALL
SELECT
c.checkpoint_id,
c.parent_checkpoint_id,
cb.version,
cb.type,
cb.blob
FROM chain prev
JOIN checkpoints c ON c.checkpoint_id = prev.parent_checkpoint_id
JOIN checkpoint_blobs cb
ON cb.thread_id = %s
AND cb.checkpoint_ns = %s
AND cb.channel = %s
AND cb.version = (c.checkpoint->'channel_versions'->>%s)::text
WHERE prev.parent_checkpoint_id IS NOT NULL
AND prev.type = 'delta'
)
SELECT DISTINCT ON (version) type, blob
FROM chain
ORDER BY version ASC
""",
(
thread_id,
checkpoint_ns,
channel,
channel,
thread_id,
checkpoint_ns,
checkpoint_id,
thread_id,
checkpoint_ns,
channel,
channel,
),
)
rows = cur.fetchall()
base = None
deltas: list[list[Any]] = []
for row in rows:
blob = self.serde.loads_typed((row["type"], row["blob"]))
if isinstance(blob, DeltaValue):
deltas.append(blob.delta)
else:
base = blob
return DeltaChainValue(base=base, deltas=deltas)
def _dump_blobs(
self,
@@ -371,3 +371,47 @@ async def test_get_checkpoint_no_channel_values(
checkpoint = await saver.aget_tuple(config)
assert checkpoint.checkpoint["channel_values"] == {}
@pytest.mark.parametrize("saver_name", ["base", "pool", "pipe"])
async def test_delta_channel_chain_reconstruction(saver_name: str) -> None:
"""AsyncPostgresSaver reconstructs DeltaChannel chain via point-lookup traversal."""
pytest.importorskip(
"langgraph.channels.delta", reason="langgraph core not installed"
)
from typing import Annotated
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
from typing_extensions import TypedDict
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
def respond(state: State) -> dict:
n = len(state["messages"])
return {"messages": [AIMessage(content=f"reply-{n}", id=f"ai-{n}")]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
async with _saver(saver_name) as saver:
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "diff-channel-test-1"}}
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="h1")]}, config)
await graph.ainvoke(
{"messages": [HumanMessage(content="there", id="h2")]}, config
)
state = await graph.aget_state(config)
msgs = state.values["messages"]
assert len(msgs) == 4, f"expected 4, got {len(msgs)}: {msgs}"
assert msgs[0].content == "hi"
assert msgs[1].content == "reply-1"
assert msgs[2].content == "there"
assert msgs[3].content == "reply-3"
@@ -1,6 +1,7 @@
from __future__ import annotations
import copy
import dataclasses
import logging
from collections.abc import AsyncIterator, Collection, Iterator, Mapping, Sequence
from typing import ( # noqa: UP035
@@ -28,6 +29,23 @@ from langgraph.checkpoint.serde.types import (
V = TypeVar("V", int, float, str)
PendingWrite = tuple[str, str, Any]
@dataclasses.dataclass
class DeltaValue:
"""Returned by DeltaChannel.checkpoint(). Represents one step's writes."""
delta: list[Any]
@dataclasses.dataclass
class DeltaChainValue:
"""Passed to DeltaChannel.from_checkpoint(). Assembled by the pregel layer."""
base: list[Any] | None # starting accumulated value; None = start from empty
deltas: list[list[Any]] # per-step write-sets, ordered oldest → newest
logger = logging.getLogger(__name__)
@@ -20,6 +20,8 @@ from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
DeltaChainValue,
DeltaValue,
SerializerProtocol,
get_checkpoint_id,
get_checkpoint_metadata,
@@ -121,17 +123,73 @@ class InMemorySaver(
return self.stack.__exit__(__exc_type, __exc_value, __traceback)
def _load_blobs(
self, thread_id: str, checkpoint_ns: str, versions: ChannelVersions
self,
thread_id: str,
checkpoint_ns: str,
versions: ChannelVersions,
checkpoint_id: str = "",
) -> dict[str, Any]:
channel_values: dict[str, Any] = {}
delta_channels: list[str] = []
for k, v in versions.items():
kk = (thread_id, checkpoint_ns, k, v)
if kk in self.blobs:
vv = self.blobs[kk]
if vv[0] != "empty":
channel_values[k] = self.serde.loads_typed(vv)
if kk not in self.blobs:
continue
vv = self.blobs[kk]
if vv[0] == "delta":
delta_channels.append(k)
elif vv[0] != "empty":
channel_values[k] = self.serde.loads_typed(vv)
for channel in delta_channels:
channel_values[channel] = self._assemble_delta_chain(
thread_id, checkpoint_ns, checkpoint_id, channel
)
return channel_values
def _assemble_delta_chain(
self,
thread_id: str,
checkpoint_ns: str,
checkpoint_id: str,
channel: str,
) -> DeltaChainValue:
"""Walk the checkpoint parent tree to collect all delta blobs for a channel."""
ns_storage = self.storage.get(thread_id, {}).get(checkpoint_ns, {})
blobs: list[Any] = []
current_id: str | None = checkpoint_id
seen_versions: set[str] = set()
while current_id is not None:
entry = ns_storage.get(current_id)
if entry is None:
break
checkpoint = self.serde.loads_typed(entry[0])
version = checkpoint["channel_versions"].get(channel)
if version is None:
break # channel not yet in this checkpoint
_, _, current_id = entry # advance to parent before the continue/break
if version in seen_versions:
continue # same blob already collected; keep walking to older checkpoints
seen_versions.add(version)
kk = (thread_id, checkpoint_ns, channel, version)
if kk not in self.blobs:
break
vv = self.blobs[kk]
if vv[0] == "empty":
break
blob = self.serde.loads_typed(vv)
blobs.append(blob)
if not isinstance(blob, DeltaValue):
break # hit a snapshot (plain list) — chain root found
blobs.reverse()
base = None
deltas: list[list[Any]] = []
for blob in blobs:
if isinstance(blob, DeltaValue):
deltas.append(blob.delta)
else:
base = blob
return DeltaChainValue(base=base, deltas=deltas)
def get_tuple(self, config: RunnableConfig) -> CheckpointTuple | None:
"""Get a checkpoint tuple from the in-memory storage.
@@ -158,7 +216,10 @@ class InMemorySaver(
checkpoint={
**checkpoint_,
"channel_values": self._load_blobs(
thread_id, checkpoint_ns, checkpoint_["channel_versions"]
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
checkpoint_id,
),
},
metadata=self.serde.loads_typed(metadata),
@@ -194,7 +255,10 @@ class InMemorySaver(
checkpoint={
**checkpoint_,
"channel_values": self._load_blobs(
thread_id, checkpoint_ns, checkpoint_["channel_versions"]
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
checkpoint_id,
),
},
metadata=self.serde.loads_typed(metadata),
@@ -304,6 +368,7 @@ class InMemorySaver(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
checkpoint_id,
),
},
metadata=metadata,
@@ -80,6 +80,8 @@ SAFE_MSGPACK_TYPES: frozenset[tuple[str, ...]] = frozenset(
("langgraph.types", "Overwrite"),
("langgraph.store.base", "Item"),
("langgraph.store.base", "GetOp"),
# DeltaChannel checkpoint value type
("langgraph.checkpoint.base", "DeltaValue"),
}
)
@@ -47,6 +47,12 @@ EMPTY_BYTES = b""
logger = logging.getLogger(__name__)
def _is_delta_value(obj: Any) -> bool:
from langgraph.checkpoint.base import DeltaValue # lazy import avoids circular dep
return isinstance(obj, DeltaValue)
class JsonPlusSerializer(SerializerProtocol):
"""Serializer that uses ormsgpack, with optional fallbacks.
@@ -239,6 +245,8 @@ class JsonPlusSerializer(SerializerProtocol):
return "bytes", obj
elif isinstance(obj, bytearray):
return "bytearray", obj
elif _is_delta_value(obj):
return "delta", _msgpack_enc({"d": obj.delta})
else:
try:
return "msgpack", _msgpack_enc(obj)
@@ -261,6 +269,13 @@ class JsonPlusSerializer(SerializerProtocol):
return ormsgpack.unpackb(
data_, ext_hook=self._unpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
elif type_ == "delta":
from langgraph.checkpoint.base import DeltaValue # lazy import
raw = ormsgpack.unpackb(
data_, ext_hook=self._unpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
return DeltaValue(delta=raw["d"])
elif self.pickle_fallback and type_ == "pickle":
return pickle.loads(data_)
else:
+13
View File
@@ -983,3 +983,16 @@ def test_msgpack_nested_pydantic_serializes_as_dict(
# No blocking should occur - inner is serialized as dict, not ext
assert "blocked" not in caplog.text.lower()
assert result == obj
def test_delta_value_serde_round_trip() -> None:
from langgraph.checkpoint.base import DeltaValue
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
serde = JsonPlusSerializer()
original = DeltaValue(delta=[{"type": "human", "content": "hi"}])
type_tag, blob = serde.dumps_typed(original)
assert type_tag == "delta"
loaded = serde.loads_typed((type_tag, blob))
assert isinstance(loaded, DeltaValue)
assert loaded.delta == original.delta
+68
View File
@@ -308,3 +308,71 @@ def test_memory_saver_with_allowlist_proxy_isolated() -> None:
assert direct is not None
expected = obj.model_dump() if hasattr(obj, "model_dump") else obj.dict()
assert direct.checkpoint["channel_values"]["foo"] == expected
class TestInMemorySaverDeltaChannel:
def test_load_blobs_assembles_delta_chain(self) -> None:
"""_load_blobs returns DeltaChainValue for delta channels, not raw DeltaValue."""
from langgraph.checkpoint.base import (
DeltaChainValue,
DeltaValue,
empty_checkpoint,
)
saver = InMemorySaver()
serde = JsonPlusSerializer()
thread_id, ns, channel = "t1", "", "messages"
v1 = "00000000000000000000000000000001.0000000000000000"
v2 = "00000000000000000000000000000002.0000000000000000"
delta1 = DeltaValue(delta=[{"content": "hi"}])
delta2 = DeltaValue(delta=[{"content": "bye"}])
saver.blobs[(thread_id, ns, channel, v1)] = serde.dumps_typed(delta1)
saver.blobs[(thread_id, ns, channel, v2)] = serde.dumps_typed(delta2)
cp1 = empty_checkpoint()
cp1["id"] = "cp1"
cp1["channel_versions"][channel] = v1
cp2 = empty_checkpoint()
cp2["id"] = "cp2"
cp2["channel_versions"][channel] = v2
saver.storage[thread_id][ns] = {
"cp1": (serde.dumps_typed(cp1), serde.dumps_typed({}), None),
"cp2": (serde.dumps_typed(cp2), serde.dumps_typed({}), "cp1"),
}
result = saver._load_blobs(thread_id, ns, {channel: v2}, "cp2")
assert channel in result
chain = result[channel]
assert isinstance(chain, DeltaChainValue)
assert chain.deltas == [[{"content": "hi"}], [{"content": "bye"}]]
def test_load_blobs_single_delta_no_parent(self) -> None:
"""Single delta with no parent checkpoint produces a chain with one delta."""
from langgraph.checkpoint.base import (
DeltaChainValue,
DeltaValue,
empty_checkpoint,
)
saver = InMemorySaver()
serde = JsonPlusSerializer()
thread_id, ns, channel = "t1", "", "messages"
v1 = "00000000000000000000000000000001.0000000000000000"
delta = DeltaValue(delta=[{"content": "only"}])
saver.blobs[(thread_id, ns, channel, v1)] = serde.dumps_typed(delta)
cp1 = empty_checkpoint()
cp1["id"] = "cp1"
cp1["channel_versions"][channel] = v1
saver.storage[thread_id][ns] = {
"cp1": (serde.dumps_typed(cp1), serde.dumps_typed({}), None)
}
result = saver._load_blobs(thread_id, ns, {channel: v1}, "cp1")
chain = result[channel]
assert isinstance(chain, DeltaChainValue)
assert chain.base is None
assert chain.deltas == [[{"content": "only"}]]
@@ -1,6 +1,7 @@
from langgraph.channels.any_value import AnyValue
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue, LastValueAfterFinish
from langgraph.channels.named_barrier_value import (
@@ -20,6 +21,7 @@ __all__ = (
"UntrackedValue",
"EphemeralValue",
"BinaryOperatorAggregate",
"DeltaChannel",
"NamedBarrierValue",
"NamedBarrierValueAfterFinish",
# topics
@@ -119,3 +119,12 @@ class BaseChannel(Generic[Value, Update, Checkpoint], ABC):
Returns `True` if the channel was updated, `False` otherwise.
"""
return False
def after_checkpoint(self, version: Any, checkpoint_id: str | None = None) -> None:
"""Called after checkpoint() with the assigned version, and after
from_checkpoint() with the current channel version.
No-op by default. Override in channels that track their own version
for incremental checkpointing (e.g. DeltaChannel).
"""
pass
+196
View File
@@ -0,0 +1,196 @@
from __future__ import annotations
from collections.abc import Callable, Sequence
from typing import Any, Generic
from langgraph.checkpoint.base import DeltaChainValue, DeltaValue
from typing_extensions import Self
from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel, Value
from langgraph.channels.binop import _get_overwrite
from langgraph.errors import EmptyChannelError
__all__ = ("DeltaChannel",)
class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]):
"""A channel that stores only per-step write deltas in checkpoints.
Reconstructs the full accumulated list at load time by replaying the
chain of deltas through the operator. Use with append-style reducers
(e.g. `add_messages`) on long-running threads to reduce checkpoint
storage from O(N²) to O(N).
Works with all checkpointers. Savers with a dedicated blob store
(InMemorySaver, PostgresSaver) use an O(1) fast-path per chain step;
all others (SQLite, MongoDB, etc.) fall back to get_tuple traversal.
Use `snapshot_every=N` to cap chain traversal depth at N steps. Every N
steps a full snapshot is written as the chain root; subsequent deltas
chain back to it, so `get_state` / reload never traverses more than N
checkpoints regardless of thread length. Recommended for savers without
a dedicated blob store.
Usage::
class State(TypedDict):
messages: Annotated[list[AnyMessage], DeltaChannel(add_messages)]
# Cap reconstruction depth (recommended for SQLite / MongoDB savers):
messages: Annotated[list[AnyMessage], DeltaChannel(add_messages, snapshot_every=50)]
# Dict-type reducer (type inferred from the Annotated outer type):
files: Annotated[dict, DeltaChannel(merge_files)]
"""
__slots__ = (
"value",
"operator",
"snapshot_every",
"_pending",
"_base_version",
"_overwritten",
"_steps_since_snapshot",
)
def __init__(
self,
operator: Callable[[list[Value], Any], list[Value]],
*,
snapshot_every: int | None = None,
) -> None:
super().__init__(list)
self.operator = operator
self.snapshot_every = snapshot_every
self.value: list[Value] = []
self._pending: list[Any] = []
self._base_version: str | None = None
self._overwritten: bool = False
self._steps_since_snapshot: int = 0
def __eq__(self, other: object) -> bool:
if not isinstance(other, DeltaChannel):
return False
if self.snapshot_every != other.snapshot_every:
return False
if (
self.operator.__name__ != "<lambda>"
and other.operator.__name__ != "<lambda>"
):
return self.operator is other.operator
return True
@property
def ValueType(self) -> Any:
return list[self.typ] # type: ignore[name-defined]
@property
def UpdateType(self) -> Any:
return self.typ | list[self.typ] # type: ignore[name-defined]
def copy(self) -> Self:
new = DeltaChannel(self.operator, snapshot_every=self.snapshot_every)
new.typ = self.typ
new.key = self.key
new.value = self.value if self.value is MISSING else self.value.copy()
new._pending = self._pending[:]
new._base_version = self._base_version
new._overwritten = self._overwritten
new._steps_since_snapshot = self._steps_since_snapshot
return new
def from_checkpoint(self, checkpoint: Any) -> Self:
new = DeltaChannel(self.operator, snapshot_every=self.snapshot_every)
new.typ = self.typ
new.key = self.key
if checkpoint is MISSING:
try:
new.value = new.typ()
except Exception:
new.value = []
elif isinstance(checkpoint, DeltaChainValue):
accumulated: list[Value] = (
checkpoint.base if checkpoint.base is not None else new.typ()
)
for step_writes in checkpoint.deltas:
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_snapshot = len(checkpoint.deltas)
elif isinstance(checkpoint, DeltaValue):
raise ValueError(
f"Channel '{self.key}' uses DeltaChannel but the checkpointer "
"does not support incremental channel storage. "
"Use InMemorySaver or PostgresSaver, or remove DeltaChannel from your schema."
)
else:
# Backwards compat: plain list from old BinaryOperatorAggregate checkpoint.
new.value = list(checkpoint)
new._pending = []
new._base_version = None # set by the subsequent after_checkpoint() call
new._overwritten = False
return new
def update(self, values: Sequence[Any]) -> bool:
if not values:
return False
seen_overwrite = False
for value in values:
is_overwrite, overwrite_value = _get_overwrite(value)
if is_overwrite:
if seen_overwrite:
from langgraph.errors import (
ErrorCode,
InvalidUpdateError,
create_error_message,
)
msg = create_error_message(
message="Can receive only one Overwrite value per super-step.",
error_code=ErrorCode.INVALID_CONCURRENT_GRAPH_UPDATE,
)
raise InvalidUpdateError(msg)
self.value = (
list(overwrite_value) if overwrite_value is not None else self.typ()
)
self._pending = list(self.value)
self._overwritten = True
seen_overwrite = True
elif not seen_overwrite:
base = self.typ() if self.value is MISSING else self.value
self.value = self.operator(base, value)
self._pending.append(value)
return True
def get(self) -> list[Value]:
if self.value is MISSING:
raise EmptyChannelError()
return self.value
def is_available(self) -> bool:
return self.value is not MISSING
def checkpoint(self) -> Any:
if (
self.snapshot_every is not None
and self._steps_since_snapshot >= self.snapshot_every
):
# Emit a full snapshot to cap chain depth at snapshot_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 DeltaValue(delta=self._pending[:])
def after_checkpoint(self, version: Any, checkpoint_id: str | None = None) -> None:
if version != self._base_version:
if self._base_version is None:
pass # First call after from_checkpoint — anchor without counting a step.
elif self.snapshot_every is not None:
if self._steps_since_snapshot >= self.snapshot_every:
self._steps_since_snapshot = 0
else:
self._steps_since_snapshot += 1
self._base_version = version
self._pending = []
self._overwritten = False
+16 -1
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import collections.abc
import inspect
import logging
import typing
@@ -47,7 +48,8 @@ from langgraph._internal._pydantic import create_model
from langgraph._internal._runnable import coerce_to_runnable
from langgraph._internal._typing import EMPTY_SEQ, MISSING, DeprecatedKwargs
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.binop import BinaryOperatorAggregate, _strip_extras
from langgraph.channels.delta import DeltaChannel
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue, LastValueAfterFinish
from langgraph.channels.named_barrier_value import (
@@ -1082,6 +1084,7 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
CompiledStateGraph: The compiled `StateGraph`.
"""
checkpointer = ensure_valid_checkpointer(checkpointer)
serde_allowlist: set[tuple[str, ...]] | None = None
if _serde.STRICT_MSGPACK_ENABLED:
schema_types: list[type[Any]] = [
@@ -1667,6 +1670,18 @@ def _is_field_channel(typ: type[Any]) -> BaseChannel | None:
# Search through all annotated medata to find channel annotations
for item in meta:
if isinstance(item, BaseChannel):
if isinstance(item, DeltaChannel) and hasattr(typ, "__origin__"):
outer = _strip_extras(typ.__origin__)
if outer in (
collections.abc.Sequence,
collections.abc.MutableSequence,
):
outer = list
item.typ = outer
try:
item.value = outer()
except Exception:
item.value = []
return item
elif isclass(item) and issubclass(item, BaseChannel):
# ex, Annotated[int, EphemeralValue, SomeOtherAnnotation]
@@ -1,5 +1,6 @@
from __future__ import annotations
import logging
from collections.abc import Mapping
from datetime import datetime, timezone
@@ -12,6 +13,8 @@ from langgraph.managed.base import ManagedValueMapping, ManagedValueSpec
LATEST_VERSION = 4
logger = logging.getLogger(__name__)
def empty_checkpoint() -> Checkpoint:
return Checkpoint(
@@ -67,13 +70,12 @@ def channels_from_checkpoint(
channel_specs[k] = v
else:
managed_specs[k] = v
return (
{
k: v.from_checkpoint(checkpoint["channel_values"].get(k, MISSING))
for k, v in channel_specs.items()
},
managed_specs,
)
channels: dict[str, BaseChannel] = {}
for k, v in channel_specs.items():
ch = v.from_checkpoint(checkpoint["channel_values"].get(k, MISSING))
ch.after_checkpoint(checkpoint["channel_versions"].get(k), checkpoint.get("id"))
channels[k] = ch
return channels, managed_specs
def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
+6
View File
@@ -881,6 +881,12 @@ class PregelLoop:
id=self.checkpoint["id"] if exiting else None,
updated_channels=self.updated_channels,
)
if do_checkpoint and self.channels:
for k, ch in self.channels.items():
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()
+10 -10
View File
@@ -1049,13 +1049,14 @@ class Pregel(
step = saved.metadata.get("step", -1) + 1
stop = step + 2
checkpoint = saved.checkpoint
channels, managed = channels_from_checkpoint(
self.channels,
saved.checkpoint,
checkpoint,
)
# tasks for this checkpoint
next_tasks = prepare_next_tasks(
saved.checkpoint,
checkpoint,
saved.pending_writes or [],
self.nodes,
channels,
@@ -1168,13 +1169,14 @@ class Pregel(
step = saved.metadata.get("step", -1) + 1
stop = step + 2
checkpoint = saved.checkpoint
channels, managed = channels_from_checkpoint(
self.channels,
saved.checkpoint,
checkpoint,
)
# tasks for this checkpoint
next_tasks = prepare_next_tasks(
saved.checkpoint,
checkpoint,
saved.pending_writes or [],
self.nodes,
channels,
@@ -1520,9 +1522,8 @@ class Pregel(
saved = checkpointer.get_tuple(config)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
)
base_checkpoint = saved.checkpoint if saved else empty_checkpoint()
checkpoint = copy_checkpoint(base_checkpoint) if saved else base_checkpoint
checkpoint_previous_versions = (
saved.checkpoint["channel_versions"].copy() if saved else {}
)
@@ -1966,9 +1967,8 @@ class Pregel(
saved = await checkpointer.aget_tuple(config)
if saved is not None:
self._migrate_checkpoint(saved.checkpoint)
checkpoint = (
copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
)
base_checkpoint = saved.checkpoint if saved else empty_checkpoint()
checkpoint = copy_checkpoint(base_checkpoint) if saved else base_checkpoint
checkpoint_previous_versions = (
saved.checkpoint["channel_versions"].copy() if saved else {}
)
+411
View File
@@ -117,3 +117,414 @@ def test_untracked_value() -> None:
new_channel = UntrackedValue(dict).from_checkpoint(checkpoint)
with pytest.raises(EmptyChannelError):
new_channel.get()
def test_delta_channel_basic_two_steps() -> None:
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.base import DeltaValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
ch = DeltaChannel(add_messages).from_checkpoint(MISSING)
ch.after_checkpoint(None)
# Step 1: one message added
ch.update([HumanMessage(content="hi", id="h1")])
d1 = ch.checkpoint()
assert isinstance(d1, DeltaValue)
assert len(d1.delta) == 1
ch.after_checkpoint("v1", checkpoint_id="cid1")
# Step 2: another message
ch.update([AIMessage(content="hello", id="a1")])
d2 = ch.checkpoint()
assert len(d2.delta) == 1
ch.after_checkpoint("v2")
# Full accumulated value is preserved in memory
assert len(ch.get()) == 2
assert ch.get()[0].content == "hi"
assert ch.get()[1].content == "hello"
def test_delta_channel_after_checkpoint_no_op_when_unchanged() -> None:
from langchain_core.messages import HumanMessage
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
ch = DeltaChannel(add_messages).from_checkpoint(MISSING)
ch.after_checkpoint(None)
ch.update([HumanMessage(content="hi", id="h1")])
ch.after_checkpoint("v1")
# Same version: no-op
ch.after_checkpoint("v1")
assert ch._base_version == "v1"
assert ch._pending == []
def test_delta_channel_from_checkpoint_chain() -> None:
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.base import DeltaChainValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
spec = DeltaChannel(add_messages)
chain = DeltaChainValue(
base=None,
deltas=[
[HumanMessage(content="hi", id="h1")],
[AIMessage(content="hello", id="a1")],
[HumanMessage(content="bye", id="h2")],
],
)
ch = spec.from_checkpoint(chain)
msgs = ch.get()
assert len(msgs) == 3
assert msgs[0].content == "hi"
assert msgs[1].content == "hello"
assert msgs[2].content == "bye"
def test_delta_channel_from_checkpoint_backwards_compat() -> None:
from langchain_core.messages import HumanMessage
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
# Old BinaryOperatorAggregate checkpoint: plain list
spec = DeltaChannel(add_messages)
old_value = [HumanMessage(content="old", id="h1")]
ch = spec.from_checkpoint(old_value)
assert ch.get() == old_value
def test_delta_channel_overwrite_resets_chain() -> None:
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.base import DeltaValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
from langgraph.types import Overwrite
ch = DeltaChannel(add_messages).from_checkpoint(MISSING)
ch.after_checkpoint(None)
ch.update([HumanMessage(content="old", id="h1")])
ch.after_checkpoint("v1")
ch.update([Overwrite([HumanMessage(content="new", id="h2")])])
d = ch.checkpoint()
assert isinstance(d, DeltaValue)
assert len(d.delta) == 1
assert d.delta[0].content == "new"
# _overwritten flag must be set so next checkpoint acts as a chain root
assert ch._overwritten is True
def test_delta_channel_unsupported_saver_raises() -> None:
"""from_checkpoint raises ValueError when the saver returns a raw DeltaValue."""
from langgraph.checkpoint.base import DeltaValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
spec = DeltaChannel(add_messages)
raw = DeltaValue(delta=[{"type": "human", "content": "hello"}])
with pytest.raises(
ValueError, match="does not support incremental channel storage"
):
spec.from_checkpoint(raw)
def test_delta_channel_remove_message_delta_and_replay() -> None:
"""RemoveMessage stored in a delta must round-trip correctly through the chain."""
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
from langgraph.checkpoint.base import DeltaChainValue, DeltaValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
spec = DeltaChannel(add_messages)
ch = spec.from_checkpoint(MISSING)
ch.after_checkpoint(None)
# Step 1: add two messages
ch.update([HumanMessage(content="hi", id="h1")])
ch.update([AIMessage(content="hello", id="a1")])
d1 = ch.checkpoint()
assert isinstance(d1, DeltaValue)
ch.after_checkpoint("v1", checkpoint_id="cid1")
assert ch.get() == [
HumanMessage(content="hi", id="h1"),
AIMessage(content="hello", id="a1"),
]
# Step 2: remove the AI message
ch.update([RemoveMessage(id="a1")])
d2 = ch.checkpoint()
assert isinstance(d2, DeltaValue)
assert any(isinstance(w, RemoveMessage) for w in d2.delta)
ch.after_checkpoint("v2", checkpoint_id="cid2")
assert ch.get() == [HumanMessage(content="hi", id="h1")]
# Replay the full chain from scratch — must reproduce the post-remove state
chain = DeltaChainValue(base=None, deltas=[d1.delta, d2.delta])
ch2 = spec.from_checkpoint(chain)
assert ch2.get() == [HumanMessage(content="hi", id="h1")]
def test_delta_channel_update_by_id_delta_and_replay() -> None:
"""Updating a message by ID stored in a delta must round-trip correctly."""
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.base import DeltaChainValue, DeltaValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
spec = DeltaChannel(add_messages)
ch = spec.from_checkpoint(MISSING)
ch.after_checkpoint(None)
# Step 1: add a message
ch.update([HumanMessage(content="original", id="h1")])
d1 = ch.checkpoint()
assert isinstance(d1, DeltaValue)
ch.after_checkpoint("v1", checkpoint_id="cid1")
# Step 2: update the same message by ID
ch.update([HumanMessage(content="updated", id="h1")])
d2 = ch.checkpoint()
assert isinstance(d2, DeltaValue)
ch.after_checkpoint("v2", checkpoint_id="cid2")
assert ch.get() == [HumanMessage(content="updated", id="h1")]
# Replay the full chain — must produce the updated message, not the original
chain = DeltaChainValue(base=None, deltas=[d1.delta, d2.delta])
ch2 = spec.from_checkpoint(chain)
assert len(ch2.get()) == 1
assert ch2.get()[0].content == "updated"
def test_delta_channel_snapshot_every_emits_plain_list() -> None:
"""snapshot_every=N causes a plain-list snapshot after N steps; next deltas chain to it."""
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.base import DeltaValue
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
SNAP = 3
spec = DeltaChannel(add_messages, snapshot_every=SNAP)
ch = spec.from_checkpoint(MISSING)
# First after_checkpoint anchors _base_version without counting a step.
ch.after_checkpoint("v0", checkpoint_id="cid0")
# Steps 1..SNAP: each should stay as DeltaValue; counter increments each step.
for i in range(1, SNAP + 1):
ch.update([HumanMessage(content=f"m{i}", id=f"h{i}")])
ckpt = ch.checkpoint()
assert isinstance(ckpt, DeltaValue), f"expected DeltaValue at step {i}"
ch.after_checkpoint(f"v{i}", checkpoint_id=f"cid{i}")
# Step SNAP+1: _steps_since_snapshot == SNAP → snapshot fires
ch.update([HumanMessage(content="snap", id="hsnap")])
snap = ch.checkpoint()
assert isinstance(snap, list), "expected plain-list snapshot at snapshot_every step"
assert len(snap) == SNAP + 1
# After snapshot, counter resets — next step is DeltaValue again
ch.after_checkpoint("vsnap", checkpoint_id="cidsnap")
ch.update([HumanMessage(content="post", id="hpost")])
post = ch.checkpoint()
assert isinstance(post, DeltaValue)
def test_delta_channel_snapshot_every_end_to_end() -> None:
"""Graph with snapshot_every: get_state returns correct accumulated value after snapshot."""
from typing import Annotated
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages, snapshot_every=2)]
counter = {"n": 0}
def respond(state: State) -> dict:
counter["n"] += 1
return {
"messages": [
AIMessage(content=f"ai-{counter['n']}", id=f"ai-{counter['n']}")
]
}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
graph = builder.compile(checkpointer=InMemorySaver())
config = {"configurable": {"thread_id": "snap-test"}}
# Run 5 turns — snapshot fires after 2 steps, then again after 2 more
for i in range(5):
graph.invoke({"messages": [HumanMessage(content=f"h{i}", id=f"h{i}")]}, config)
state = graph.get_state(config)
msgs = state.values["messages"]
# 5 human + 5 AI = 10 total
assert len(msgs) == 10, f"expected 10 messages, got {len(msgs)}: {msgs}"
def test_delta_channel_inmemory_saver_assembles_chain() -> None:
"""InMemorySaver assembles the delta chain inside get_tuple (no pregel involvement)."""
from typing import Annotated
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
n = {"v": 0}
def respond(state: State) -> dict:
n["v"] += 1
return {"messages": [AIMessage(content=f"ok{n['v']}", id=f"ai{n['v']}")]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
saver = InMemorySaver()
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "t1"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="h1")]}, config)
graph.invoke({"messages": [HumanMessage(content="bye", id="h2")]}, config)
# get_tuple must return a fully assembled DeltaChainValue, not raw DeltaValue
from langgraph.checkpoint.base import DeltaChainValue, DeltaValue
saved = saver.get_tuple(config)
assert saved is not None
assert "messages" in saved.checkpoint["channel_values"]
assert not isinstance(saved.checkpoint["channel_values"]["messages"], DeltaValue)
assert isinstance(saved.checkpoint["channel_values"]["messages"], DeltaChainValue)
state = graph.get_state(config)
assert len(state.values["messages"]) == 4 # 2 human + 2 AI
def _delta_channel_with_type(operator, typ):
"""Build a DeltaChannel with an explicit type via the Annotated injection path."""
from typing import Annotated
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.state import _get_channel
return _get_channel("_test", Annotated[typ, DeltaChannel(operator)])
def test_delta_channel_dict_reducer_fresh_channel() -> None:
"""DeltaChannel with a dict reducer starts as empty dict on MISSING checkpoint."""
def merge_dicts(left: dict, right: dict) -> dict:
return {**left, **right}
ch = _delta_channel_with_type(merge_dicts, dict).from_checkpoint(MISSING)
# Should be available (not raise EmptyChannelError) and start empty
assert ch.is_available()
assert ch.get() == {}
def test_delta_channel_dict_reducer_basic_updates() -> None:
"""DeltaChannel with a dict reducer accumulates key/value pairs across steps."""
from langgraph.checkpoint.base import DeltaValue
def merge_dicts(left: dict, right: dict) -> dict:
return {**left, **right}
ch = _delta_channel_with_type(merge_dicts, dict).from_checkpoint(MISSING)
ch.after_checkpoint(None)
ch.update([{"a": 1}])
d1 = ch.checkpoint()
assert isinstance(d1, DeltaValue)
assert d1.delta == [{"a": 1}]
ch.after_checkpoint("v1", checkpoint_id="cid1")
ch.update([{"b": 2}])
d2 = ch.checkpoint()
assert d2.delta == [{"b": 2}]
ch.after_checkpoint("v2")
assert ch.get() == {"a": 1, "b": 2}
def test_delta_channel_dict_reducer_chain_reconstruction() -> None:
"""DeltaChainValue replays correctly through a dict merge reducer."""
from langgraph.checkpoint.base import DeltaChainValue
def merge_dicts(left: dict, right: dict) -> dict:
return {**left, **right}
spec = _delta_channel_with_type(merge_dicts, dict)
chain = DeltaChainValue(
base={"a": 1},
deltas=[[{"b": 2}], [{"c": 3}]],
)
ch = spec.from_checkpoint(chain)
assert ch.get() == {"a": 1, "b": 2, "c": 3}
assert ch._steps_since_snapshot == 2
def test_delta_channel_dict_reducer_with_deletions() -> None:
"""Dict reducer that treats None values as deletions works end-to-end (deepagents pattern)."""
from langgraph.checkpoint.base import DeltaChainValue
def merge_files(left: dict | None, right: dict) -> dict:
if left is None:
return {k: v for k, v in right.items() if v is not None}
result = {**left}
for k, v in right.items():
if v is None:
result.pop(k, None)
else:
result[k] = v
return result
ch = _delta_channel_with_type(merge_files, dict).from_checkpoint(MISSING)
ch.after_checkpoint(None)
ch.update([{"file1.py": "content1", "file2.py": "content2"}])
ch.after_checkpoint("v1", checkpoint_id="cid1")
# Delete file1, add file3
ch.update([{"file1.py": None, "file3.py": "content3"}])
ch.after_checkpoint("v2", checkpoint_id="cid2")
assert ch.get() == {"file2.py": "content2", "file3.py": "content3"}
# Confirm chain reconstruction produces the same result
chain = DeltaChainValue(
base={},
deltas=[
[{"file1.py": "content1", "file2.py": "content2"}],
[{"file1.py": None, "file3.py": "content3"}],
],
)
spec = _delta_channel_with_type(merge_files, dict)
ch2 = spec.from_checkpoint(chain)
assert ch2.get() == {"file2.py": "content2", "file3.py": "content3"}
@@ -0,0 +1,379 @@
"""Benchmark: DeltaChannel vs BinaryOperatorAggregate storage and time.
Run directly: python tests/test_delta_channel_benchmark.py
Run via pytest: pytest tests/test_delta_channel_benchmark.py -s
Simulates realistic multi-turn conversations with paragraph-length messages
(~100 tokens each) scaling up to 1M-token-equivalent histories.
Token estimates: 1 token ≈ 4 chars; each turn ≈ 200 tokens (human + AI).
A 1M-token conversation ≈ 5,000 turns of realistic messages.
"""
from __future__ import annotations
import sys
import time
from typing import Annotated, Any
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import MemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, StateGraph
from langgraph.graph.message import add_messages
try:
from langgraph.checkpoint.sqlite import SqliteSaver
_SQLITE_AVAILABLE = True
except ImportError:
_SQLITE_AVAILABLE = False
try:
from langgraph.checkpoint.postgres import PostgresSaver
_POSTGRES_AVAILABLE = True
_POSTGRES_URI = (
"postgres://postgres:postgres@localhost:5441/postgres?sslmode=disable"
)
except ImportError:
_POSTGRES_AVAILABLE = False
SNAPSHOT_EVERY = 50
# ---------------------------------------------------------------------------
# Realistic message payload (~100 tokens / ~400 chars each)
# ---------------------------------------------------------------------------
_HUMAN_TEMPLATE = (
"I need help understanding the implications of {topic} on our system architecture. "
"Specifically, I'm concerned about how this interacts with our existing {concern} "
"and whether we need to refactor the {component} layer before proceeding."
)
_AI_TEMPLATE = (
"Great question about {topic}. The key insight here is that {concern} introduces "
"a subtle ordering dependency that most teams overlook until they hit it in production. "
"For your {component} layer specifically, I'd recommend starting with a careful audit "
"of the interface boundaries before making any structural changes. This will give you "
"a clear picture of the blast radius and let you sequence the migration safely."
)
_TOPICS = [
"distributed tracing",
"eventual consistency",
"schema migration",
"backpressure handling",
"idempotency guarantees",
"cache invalidation",
"connection pooling",
"rate limiting",
"circuit breaking",
"observability pipelines",
]
_CONCERNS = [
"concurrency model",
"retry semantics",
"state management",
"error propagation",
"latency budget",
]
_COMPONENTS = [
"persistence",
"routing",
"ingestion",
"aggregation",
"serialization",
]
def _human_content(i: int) -> str:
return _HUMAN_TEMPLATE.format(
topic=_TOPICS[i % len(_TOPICS)],
concern=_CONCERNS[i % len(_CONCERNS)],
component=_COMPONENTS[i % len(_COMPONENTS)],
)
def _ai_content(i: int) -> str:
return _AI_TEMPLATE.format(
topic=_TOPICS[i % len(_TOPICS)],
concern=_CONCERNS[i % len(_CONCERNS)],
component=_COMPONENTS[i % len(_COMPONENTS)],
)
# ---------------------------------------------------------------------------
# State definitions
# ---------------------------------------------------------------------------
class BinaryState(TypedDict):
messages: Annotated[list, add_messages]
class DeltaState(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
class DeltaSnapshotState(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages, snapshot_every=SNAPSHOT_EVERY)]
# ---------------------------------------------------------------------------
# Graph factory
# ---------------------------------------------------------------------------
def _make_graph(state_cls: type, checkpointer: Any = None) -> Any:
def human_node(state: Any) -> dict:
return {}
def ai_node(state: Any) -> dict:
i = len(state["messages"]) // 2
return {"messages": [AIMessage(content=_ai_content(i), id=f"a{i}")]}
g = StateGraph(state_cls)
g.add_node("human", human_node)
g.add_node("ai", ai_node)
g.add_edge("human", "ai")
g.add_edge("ai", END)
g.set_entry_point("human")
return g.compile(checkpointer=checkpointer or MemorySaver())
# ---------------------------------------------------------------------------
# Measurement helpers
# ---------------------------------------------------------------------------
def _total_blob_bytes(saver: MemorySaver) -> int:
total = 0
for (_, _, _, _), (type_tag, blob) in saver.blobs.items():
if blob is not None:
total += len(blob)
return total
def _run_turns(
n_turns: int,
state_cls: type,
checkpointer: Any = None,
) -> tuple[float, float, int]:
"""Run n_turns conversation turns.
Returns (write_elapsed_s, read_elapsed_s, total_blob_bytes).
blob_bytes is -1 for savers without in-memory blob stores (e.g. SQLite).
Read latency is measured as the time to invoke the graph with no new
messages after the full history is built — this forces state rehydration.
"""
graph = _make_graph(state_cls, checkpointer)
config = {"configurable": {"thread_id": "bench"}}
t0 = time.perf_counter()
for i in range(n_turns):
graph.invoke(
{"messages": [HumanMessage(content=_human_content(i), id=f"h{i}")]},
config,
)
write_elapsed = time.perf_counter() - t0
# Measure read/rehydration: get_state forces the channel to rebuild
t1 = time.perf_counter()
for _ in range(5):
graph.get_state(config)
read_elapsed = (time.perf_counter() - t1) / 5
if isinstance(graph.checkpointer, MemorySaver):
blob_bytes = _total_blob_bytes(graph.checkpointer)
else:
blob_bytes = -1
return write_elapsed, read_elapsed, blob_bytes
def _fmt_bytes(n: int) -> str:
if n >= 1_000_000:
return f"{n / 1_000_000:.1f} MB"
if n >= 1_000:
return f"{n / 1_000:.1f} KB"
return f"{n} B"
def _approx_tokens(n_turns: int) -> str:
# ~100 tokens human + ~100 tokens AI per turn
tokens = n_turns * 200
if tokens >= 1_000_000:
return f"~{tokens / 1_000_000:.1f}M tok"
if tokens >= 1_000:
return f"~{tokens / 1_000:.0f}K tok"
return f"~{tokens} tok"
# ---------------------------------------------------------------------------
# Benchmark matrix
# ---------------------------------------------------------------------------
# Turn counts chosen to demonstrate O(N²) vs O(N) storage growth without running too long.
# Extrapolation: 5,000 turns × ~200 tokens/turn ≈ 1M tokens (Claude's full context window).
TURN_COUNTS = [10, 25, 50, 100]
def _checkpointer_factories() -> list[tuple[str, Any]]:
"""Return (label, context_manager_or_none) pairs for available checkpointers."""
return [("InMemory", None)]
def run_benchmark() -> None:
print()
print(
"DeltaChannel vs add_messages (BinaryOperatorAggregate) — checkpoint storage & latency"
)
print("Simulating realistic multi-turn conversations up to ~1M-token histories")
print("(5,000 turns × ~200 tokens/turn ≈ 1M tokens — Claude's full context window)")
print()
checkpointers: list[tuple[str, Any]] = [("InMemory", None)]
if _POSTGRES_AVAILABLE:
checkpointers.append(("Postgres (recursive CTE)", "postgres"))
for cp_label, cp_hint in checkpointers:
print(f"--- Checkpointer: {cp_label} ---")
_run_benchmark_for_checkpointer(cp_hint)
def _run_benchmark_for_checkpointer(cp_hint: Any) -> None:
import contextlib
import tempfile
@contextlib.contextmanager
def _make_saver():
if cp_hint is None:
yield None
elif cp_hint == "postgres":
with PostgresSaver.from_conn_string(_POSTGRES_URI) as saver:
saver.setup()
with saver._cursor() as cur:
cur.execute("DELETE FROM checkpoints WHERE thread_id = 'bench'")
cur.execute(
"DELETE FROM checkpoint_blobs WHERE thread_id = 'bench'"
)
cur.execute(
"DELETE FROM checkpoint_writes WHERE thread_id = 'bench'"
)
yield saver
else:
with tempfile.NamedTemporaryFile(suffix=".db") as f:
with SqliteSaver.from_conn_string(f.name) as saver:
yield saver
rows = []
for turns in TURN_COUNTS:
with _make_saver() as saver:
b_wt, b_rt, b_bytes = _run_turns(turns, BinaryState, saver)
with _make_saver() as saver:
d_wt, d_rt, d_bytes = _run_turns(turns, DeltaState, saver)
with _make_saver() as saver:
s_wt, s_rt, s_bytes = _run_turns(turns, DeltaSnapshotState, saver)
rows.append((turns, b_bytes, d_bytes, s_bytes, b_rt, d_rt, s_rt))
# ── Table 1: Storage ─────────────────────────────────────────────────────
W = 80
print("Storage (checkpoint blob bytes)")
print("=" * W)
print(
f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12} {'delta+snap':>12} {'savings':>8}"
)
print("-" * W)
storage_results = []
for turns, b_bytes, d_bytes, s_bytes, *_ in rows:
if b_bytes < 0:
print(
f"{turns:>6} {_approx_tokens(turns):>10} {'n/a':>12} {'n/a':>12} {'n/a':>12} {'n/a':>8}"
)
else:
ratio = b_bytes / s_bytes if s_bytes else float("inf")
storage_results.append((turns, b_bytes, s_bytes, ratio))
print(
f"{turns:>6} {_approx_tokens(turns):>10} "
f"{_fmt_bytes(b_bytes):>12} {_fmt_bytes(d_bytes):>12} {_fmt_bytes(s_bytes):>12} "
f"{ratio:>7.0f}x"
)
print("=" * W)
print()
# ── Table 2: Read latency ─────────────────────────────────────────────────
print("Read latency (avg of 5 get_state calls)")
print("=" * W)
print(
f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12} {'delta+snap':>12}"
)
print("-" * W)
for turns, b_bytes, d_bytes, s_bytes, b_rt, d_rt, s_rt in rows:
print(
f"{turns:>6} {_approx_tokens(turns):>10} "
f"{b_rt * 1000:>10.1f}ms {d_rt * 1000:>10.1f}ms {s_rt * 1000:>10.1f}ms"
)
print("=" * W)
print()
if storage_results:
best = storage_results[-1]
turns, b_bytes, s_bytes, ratio = best
_, _, _, _, b_rt, _, s_rt = rows[-1]
print(
f"At {turns} turns: {_fmt_bytes(b_bytes)} → {_fmt_bytes(s_bytes)} ({ratio:.0f}x less storage); "
f"read {b_rt * 1000:.1f}ms → {s_rt * 1000:.1f}ms"
)
print()
print("Legend:")
print(" add_msgs = Annotated[list, add_messages] — O(N²) storage")
print(
" delta = DeltaChannel(add_messages) — O(N) storage, unbounded chain"
)
print(
f" delta+snap = DeltaChannel(add_messages, snapshot_every={SNAPSHOT_EVERY}) — O(N) storage, bounded read depth"
)
print()
# ---------------------------------------------------------------------------
# Pytest entry point
# ---------------------------------------------------------------------------
@pytest.mark.skip(
reason="slow benchmark — run manually with: python tests/test_delta_channel_benchmark.py"
)
def test_delta_channel_benchmark(capsys: Any) -> None:
"""Storage grows O(N²) for add_messages, O(N) for DeltaChannel."""
with capsys.disabled():
run_benchmark()
# Correctness assertion: DeltaChannel must use less storage at scale.
for turns in [25, 50]:
_, _, b_bytes = _run_turns(turns, BinaryState)
_, _, d_bytes = _run_turns(turns, DeltaState)
_, _, s_bytes = _run_turns(turns, DeltaSnapshotState)
assert d_bytes < b_bytes, (
f"DeltaChannel should use less storage at {turns} turns, "
f"got delta={d_bytes} binary={b_bytes}"
)
assert s_bytes < b_bytes, (
f"DeltaChannel+snapshot should use less storage at {turns} turns, "
f"got snapshot={s_bytes} binary={b_bytes}"
)
# ---------------------------------------------------------------------------
# Script entry point
# ---------------------------------------------------------------------------
if __name__ == "__main__":
run_benchmark()
sys.exit(0)
+187
View File
@@ -9400,3 +9400,190 @@ def test_fork_does_not_apply_pending_writes(
# Should be: 1 (input) + 20 (forked node_a) + 100 (node_b) = 121
assert result == {"value": 121}
async def test_delta_channel_end_to_end_inmemory() -> None:
"""Full graph run: DeltaChannel accumulates correctly across multiple turns."""
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
def respond(state: State) -> dict:
n = len(state["messages"])
return {"messages": [AIMessage(content=f"reply-{n}", id=f"ai-{n}")]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
graph = builder.compile(checkpointer=InMemorySaver())
config = {"configurable": {"thread_id": "diff-test-1"}}
# Turn 1
graph.invoke({"messages": [HumanMessage(content="hello", id="h1")]}, config)
# Turn 2
graph.invoke({"messages": [HumanMessage(content="world", id="h2")]}, config)
# Turn 3
graph.invoke({"messages": [HumanMessage(content="bye", id="h3")]}, config)
state = graph.get_state(config)
msgs = state.values["messages"]
# 3 human + 3 AI = 6 total
assert len(msgs) == 6, f"expected 6 messages, got {len(msgs)}: {msgs}"
assert msgs[0].content == "hello"
assert msgs[2].content == "world"
assert msgs[4].content == "bye"
assert msgs[1].content == "reply-1"
assert msgs[3].content == "reply-3"
assert msgs[5].content == "reply-5"
async def test_delta_channel_time_travel() -> None:
"""Time-travel back to turn-1 checkpoint and resume; continuation must not include turn-2 deltas."""
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
counter = {"n": 0}
def respond(state: State) -> dict:
counter["n"] += 1
return {
"messages": [
AIMessage(content=f"ai-{counter['n']}", id=f"ai-{counter['n']}")
]
}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_edge(START, "respond")
saver = InMemorySaver()
graph = builder.compile(checkpointer=saver)
config = {"configurable": {"thread_id": "diff-time-travel"}}
# Run 2 turns: h1→ai-1, h2→ai-2
graph.invoke({"messages": [HumanMessage(content="h1", id="h1")]}, config)
graph.invoke({"messages": [HumanMessage(content="h2", id="h2")]}, config)
# Find the checkpoint after turn 1 (2 messages: h1 + ai-1)
history = list(graph.get_state_history(config))
after_turn1 = next(h for h in history if len(h.values.get("messages", [])) == 2)
assert len(after_turn1.values["messages"]) == 2
assert after_turn1.values["messages"][0].content == "h1"
assert after_turn1.values["messages"][1].content == "ai-1"
# Resume from turn-1 checkpoint: inject h3, expect 3 messages total (h1, ai-1, ai-N)
# NOT 5 messages (turn-2 deltas must not bleed into the resumed run)
result = graph.invoke(
{"messages": [HumanMessage(content="h3", id="h3")]},
after_turn1.config,
)
msgs = result["messages"]
# Should be: h1, ai-1, h3, ai-N — 4 messages total
assert len(msgs) == 4, (
f"expected 4 messages after time-travel resume, got {len(msgs)}: {msgs}"
)
assert msgs[0].content == "h1"
assert msgs[1].content == "ai-1"
assert msgs[2].content == "h3"
async def test_delta_channel_remove_message_end_to_end() -> None:
"""RemoveMessage inside a DeltaChannel graph must persist and reload correctly."""
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
def respond(state: State) -> dict:
return {"messages": [AIMessage(content="reply", id="ai-1")]}
def delete_first(state: State) -> dict:
# removes the first message
return {"messages": [RemoveMessage(id=state["messages"][0].id)]}
builder = StateGraph(State)
builder.add_node("respond", respond)
builder.add_node("delete_first", delete_first)
builder.add_edge(START, "respond")
builder.add_edge("respond", "delete_first")
graph = builder.compile(checkpointer=InMemorySaver())
config = {"configurable": {"thread_id": "diff-remove-test"}}
graph.invoke({"messages": [HumanMessage(content="hello", id="h1")]}, config)
state = graph.get_state(config)
msgs = state.values["messages"]
# h1 was removed, only ai-1 should remain
assert len(msgs) == 1, f"expected 1 message, got {len(msgs)}: {msgs}"
assert msgs[0].id == "ai-1"
# A subsequent turn must reconstruct from the checkpoint correctly
graph.invoke({"messages": [HumanMessage(content="again", id="h2")]}, config)
state = graph.get_state(config)
msgs = state.values["messages"]
# ai-1 + h2 + ai-1(second reply, same id overwrites) + h2 removed
# more simply: after second run we expect ai-1 updated + h2 remaining minus deleted h2
# just assert h1 is still gone
assert all(m.id != "h1" for m in msgs), (
"h1 should still be absent after second turn"
)
async def test_delta_channel_update_by_id_end_to_end() -> None:
"""Updating a message by ID via DeltaChannel must persist and reload correctly."""
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import START, StateGraph
from langgraph.graph.message import add_messages
class State(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
def update_msg(state: State) -> dict:
# re-send h1 with updated content
return {"messages": [HumanMessage(content="updated", id="h1")]}
builder = StateGraph(State)
builder.add_node("update_msg", update_msg)
builder.add_edge(START, "update_msg")
graph = builder.compile(checkpointer=InMemorySaver())
config = {"configurable": {"thread_id": "diff-update-id-test"}}
graph.invoke({"messages": [HumanMessage(content="original", id="h1")]}, config)
state = graph.get_state(config)
msgs = state.values["messages"]
assert len(msgs) == 1, f"expected 1 message, got {len(msgs)}: {msgs}"
assert msgs[0].content == "updated"
assert msgs[0].id == "h1"
# Second turn: verify the updated state is the base for further accumulation
graph.invoke({"messages": [HumanMessage(content="new", id="h2")]}, config)
state = graph.get_state(config)
msgs = state.values["messages"]
ids = [m.id for m in msgs]
assert "h1" in ids # h1 persists (updated, not duplicated)
assert "h2" in ids
assert ids.count("h1") == 1, "h1 must not be duplicated"
@@ -0,0 +1,185 @@
"""Sweep snapshot_every values to find the storage vs. time-travel tradeoff.
Run directly: python tests/test_rehydrate_sweep.py
Run via pytest: pytest tests/test_rehydrate_sweep.py -s
"""
from __future__ import annotations
import sys
import time
from typing import Annotated, Any
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import MemorySaver
from typing_extensions import TypedDict
from langgraph.channels.delta import DeltaChannel
from langgraph.graph import END, StateGraph
from langgraph.graph.message import add_messages
# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
REHYDRATE_SWEEP = [5, 10, 25, 50, 100, None] # None = no rehydration (pure diff)
TURN_COUNTS = [50, 100, 250, 500]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_state(snapshot_every: int | None) -> type:
channel = DeltaChannel(add_messages, snapshot_every=snapshot_every)
return TypedDict("S", {"messages": Annotated[list, channel]})
def _make_graph(state_cls: type) -> Any:
def human_node(state: Any) -> dict:
return {}
def ai_node(state: Any) -> dict:
last = state["messages"][-1]
return {"messages": [AIMessage(content=f"reply-to-{last.id}")]}
g = StateGraph(state_cls)
g.add_node("human", human_node)
g.add_node("ai", ai_node)
g.add_edge("human", "ai")
g.add_edge("ai", END)
g.set_entry_point("human")
return g.compile(checkpointer=MemorySaver())
def _total_blob_bytes(saver: MemorySaver) -> int:
total = 0
for (_, _, _, _), (type_tag, blob) in saver.blobs.items():
if blob is not None:
total += len(blob)
return total
def _measure_time_travel_ms(graph: Any, config: dict) -> float:
"""Time how long it takes to get state at the very first checkpoint (worst case)."""
history = list(graph.get_state_history(config))
if not history:
return 0.0
oldest = history[-1]
t0 = time.perf_counter()
graph.get_state(oldest.config)
return (time.perf_counter() - t0) * 1000
def _run(n_turns: int, snapshot_every: int | None) -> tuple[float, int, float]:
"""Returns (write_ms, blob_bytes, time_travel_ms)."""
state_cls = _make_state(snapshot_every)
graph = _make_graph(state_cls)
saver: MemorySaver = graph.checkpointer # type: ignore[assignment]
config = {"configurable": {"thread_id": "sweep"}}
t0 = time.perf_counter()
for i in range(n_turns):
graph.invoke(
{"messages": [HumanMessage(content=f"msg-{i}", id=f"h{i}")]}, config
)
write_ms = (time.perf_counter() - t0) * 1000
blob_bytes = _total_blob_bytes(saver)
tt_ms = _measure_time_travel_ms(graph, config)
return write_ms, blob_bytes, tt_ms
# ---------------------------------------------------------------------------
# ASCII sparkline
# ---------------------------------------------------------------------------
def _sparkline(values: list[float], width: int = 20) -> str:
bars = " ▁▂▃▄▅▆▇█"
lo, hi = min(values), max(values)
span = hi - lo or 1
chars = [bars[round((v - lo) / span * (len(bars) - 1))] for v in values]
return "".join(chars).ljust(width)
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def run_sweep() -> None:
label = {v: (str(v) if v is not None else "None(∞)") for v in REHYDRATE_SWEEP}
print()
print("snapshot_every sweep — storage vs time-travel cost")
print("=" * 90)
for turns in TURN_COUNTS:
print(f"\n--- {turns} turns ---")
col_w = 12
header = (
f"{'snapshot_every':>18} "
f"{'blob_bytes':>{col_w}} "
f"{'write_ms':>{col_w}} "
f"{'time_travel_ms':>{col_w}}"
)
print(header)
print("-" * 60)
tt_vals: list[float] = []
byte_vals: list[int] = []
write_vals: list[float] = []
rows: list[tuple] = []
for rv in REHYDRATE_SWEEP:
write_ms, blob_bytes, tt_ms = _run(turns, rv)
rows.append((rv, blob_bytes, write_ms, tt_ms))
byte_vals.append(blob_bytes)
write_vals.append(write_ms)
tt_vals.append(tt_ms)
for rv, blob_bytes, write_ms, tt_ms in rows:
print(
f"{label[rv]:>18} "
f"{blob_bytes:>{col_w},} "
f"{write_ms:>{col_w}.1f} "
f"{tt_ms:>{col_w}.2f}"
)
print()
print(
f" bytes spark: [{_sparkline(byte_vals)}] "
f"lo={min(byte_vals):,} hi={max(byte_vals):,}"
)
print(
f" time-travel spark: [{_sparkline(tt_vals)}] "
f"lo={min(tt_vals):.2f}ms hi={max(tt_vals):.2f}ms"
)
print(
f" write spark: [{_sparkline(write_vals)}] "
f"lo={min(write_vals):.1f}ms hi={max(write_vals):.1f}ms"
)
print()
print("=" * 90)
print(
"snapshot_every=None means pure diff (no snapshots) — "
"lowest storage, highest time-travel cost."
)
print(
"Lower snapshot_every = more frequent full snapshots = "
"faster time-travel, more storage."
)
print()
def test_rehydrate_sweep(capsys: Any) -> None:
with capsys.disabled():
run_sweep()
if __name__ == "__main__":
run_sweep()
sys.exit(0)