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
36 changed files with 1290 additions and 1427 deletions
@@ -4,7 +4,7 @@ import threading
from collections import defaultdict
from collections.abc import Iterator, Sequence
from contextlib import contextmanager
from typing import Any, cast
from typing import Any
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
@@ -442,22 +442,14 @@ class PostgresSaver(BasePostgresSaver):
including its configuration, metadata, parent checkpoint (if any),
and pending writes.
"""
from langgraph.checkpoint.base import DeltaChannelSentinel
channel_values = self._load_blobs(value["channel_values"])
if any(isinstance(v, DeltaChannelSentinel) for v in channel_values.values()):
cp_config = cast(
RunnableConfig,
{
"configurable": {
"thread_id": value["thread_id"],
"checkpoint_ns": value["checkpoint_ns"],
"checkpoint_id": value["checkpoint_id"],
}
},
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,
)
with self._cursor() as cur:
self._resolve_delta_channels(cp_config, channel_values, cur)
return CheckpointTuple(
{
"configurable": {
@@ -13,7 +13,8 @@ from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
DeltaChannelSentinel,
DeltaChainValue,
DeltaValue,
get_checkpoint_id,
get_serializable_checkpoint_metadata,
)
@@ -392,57 +393,80 @@ class AsyncPostgresSaver(BasePostgresSaver):
async with conn.cursor(binary=True, row_factory=dict_row) as cur:
yield cur
async def _aget_channel_writes_cur(
async def _aload_delta_chain(
self,
thread_id: str,
checkpoint_ns: str,
checkpoint_id: str,
channel: str,
cur: Any,
) -> list[Any]:
"""Async version of _get_channel_writes_cur — see sync version for rationale."""
) -> DeltaChainValue:
"""Fetch the full delta chain for a channel in one recursive CTE query (async)."""
await cur.execute(
"SELECT checkpoint_id, parent_checkpoint_id FROM checkpoints "
"WHERE thread_id = %s AND checkpoint_ns = %s",
(thread_id, checkpoint_ns),
)
parent_map: dict[str, str | None] = {
row["checkpoint_id"]: row["parent_checkpoint_id"]
for row in await cur.fetchall()
}
ancestor_ids: list[str] = []
cid: str | None = parent_map.get(checkpoint_id)
while cid is not None:
ancestor_ids.append(cid)
cid = parent_map.get(cid)
if not ancestor_ids:
return []
await cur.execute(
"SELECT checkpoint_id, type, blob FROM checkpoint_writes "
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
" AND checkpoint_id = ANY(%s) "
"ORDER BY task_id, idx",
(thread_id, checkpoint_ns, channel, ancestor_ids),
)
writes_by_cp: dict[str, list[tuple[str, bytes]]] = defaultdict(list)
for row in await cur.fetchall():
writes_by_cp[row["checkpoint_id"]].append((row["type"], row["blob"]))
result = []
for cid in reversed(ancestor_ids):
for type_tag, blob in writes_by_cp.get(cid, []):
result.append(self.serde.loads_typed((type_tag, blob)))
return result
"""
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
async def aget_channel_writes(
self, config: RunnableConfig, channel: str
) -> list[Any]:
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = config["configurable"]["checkpoint_id"]
async with self._cursor() as cur:
return await self._aget_channel_writes_cur(
thread_id, checkpoint_ns, checkpoint_id, channel, cur
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:
"""
@@ -465,14 +489,12 @@ class AsyncPostgresSaver(BasePostgresSaver):
if blob_values:
channel_values = self._load_blobs(blob_values)
delta_channels = [
ch
for ch, v in channel_values.items()
if isinstance(v, DeltaChannelSentinel)
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._aget_channel_writes_cur(
channel_values[channel] = await self._aload_delta_chain(
thread_id, checkpoint_ns, checkpoint_id, channel, cur
)
@@ -2,7 +2,6 @@ from __future__ import annotations
import random
import warnings
from collections import defaultdict
from collections.abc import Sequence
from importlib.metadata import version as get_version
from typing import Any, cast
@@ -12,7 +11,8 @@ from langgraph.checkpoint.base import (
WRITES_IDX_MAP,
BaseCheckpointSaver,
ChannelVersions,
DeltaChannelSentinel,
DeltaChainValue,
DeltaValue,
get_checkpoint_id,
)
from langgraph.checkpoint.serde.types import TASKS
@@ -188,77 +188,105 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
def _load_blobs(
self,
blob_values: Any,
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 {}
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 != "empty":
result[k.decode()] = self.serde.loads_typed((type_tag, v))
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 _resolve_delta_channels(
self,
config: RunnableConfig,
channel_values: dict[str, Any],
cur: Any,
) -> None:
for channel, value in list(channel_values.items()):
if isinstance(value, DeltaChannelSentinel):
channel_values[channel] = self._get_channel_writes_cur(
config["configurable"]["thread_id"],
config["configurable"].get("checkpoint_ns", ""),
config["configurable"]["checkpoint_id"],
channel,
cur,
)
def _get_channel_writes_cur(
def _load_delta_chain(
self,
thread_id: str,
checkpoint_ns: str,
checkpoint_id: str,
channel: str,
cur: Any,
) -> list[Any]:
"""Fetch writes for `channel` across the checkpoint ancestor chain, oldest→newest.
) -> 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
Two queries:
1. Fetch all (checkpoint_id, parent_checkpoint_id) for the thread — cheap, IDs only.
2. Walk the ancestor chain in Python, then fetch writes with a plain ANY() filter.
"""
cur.execute(
"SELECT checkpoint_id, parent_checkpoint_id FROM checkpoints "
"WHERE thread_id = %s AND checkpoint_ns = %s",
(thread_id, checkpoint_ns),
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,
),
)
parent_map: dict[str, str | None] = {
row["checkpoint_id"]: row["parent_checkpoint_id"] for row in cur.fetchall()
}
ancestor_ids: list[str] = []
cid: str | None = parent_map.get(checkpoint_id)
while cid is not None:
ancestor_ids.append(cid)
cid = parent_map.get(cid)
if not ancestor_ids:
return []
cur.execute(
"SELECT checkpoint_id, type, blob FROM checkpoint_writes "
"WHERE thread_id = %s AND checkpoint_ns = %s AND channel = %s "
" AND checkpoint_id = ANY(%s) "
"ORDER BY task_id, idx",
(thread_id, checkpoint_ns, channel, ancestor_ids),
)
writes_by_cp: dict[str, list[tuple[str, bytes]]] = defaultdict(list)
for row in cur.fetchall():
writes_by_cp[row["checkpoint_id"]].append((row["type"], row["blob"]))
result = []
for cid in reversed(ancestor_ids):
for type_tag, blob in writes_by_cp.get(cid, []):
result.append(self.serde.loads_typed((type_tag, blob)))
return result
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,
@@ -3,12 +3,10 @@ from __future__ import annotations
import copy
import dataclasses
import logging
import threading
from collections.abc import AsyncIterator, Collection, Iterator, Mapping, Sequence
from typing import ( # noqa: UP035
Any,
Generic,
List,
Literal,
NamedTuple,
TypedDict,
@@ -34,17 +32,19 @@ PendingWrite = tuple[str, str, Any]
@dataclasses.dataclass
class DeltaChannelSentinel:
"""Marker stored in checkpoint_blobs for a DeltaChannel field.
class DeltaValue:
"""Returned by DeltaChannel.checkpoint(). Represents one step's writes."""
No data is stored here — the actual per-step writes live in checkpoint_writes
and are replayed through the reducer at load time.
"""
pass
delta: list[Any]
_DELTA_RECONSTRUCTION: threading.local = threading.local()
@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__)
@@ -475,56 +475,6 @@ class BaseCheckpointSaver(Generic[V]):
"""
raise NotImplementedError
def get_channel_writes(self, config: RunnableConfig, channel: str) -> List[Any]: # noqa: UP006
"""Collect all writes for `channel` across this checkpoint's ancestry, oldest→newest.
Default implementation walks the full thread history via `list()`. Savers can
override with a more efficient query (InMemorySaver and PostgresSaver do this).
"""
# Guard against re-entrant calls: when list() triggers reconstruction which
# calls list() again, the inner call returns tuples with DeltaChannelSentinel
# in channel_values (which get_channel_writes ignores — it only reads
# pending_writes). This breaks the recursion safely.
if getattr(_DELTA_RECONSTRUCTION, "active", False):
return []
_DELTA_RECONSTRUCTION.active = True
try:
result: list[Any] = []
target_id = config["configurable"].get("checkpoint_id")
for tup in self.list(config):
if tup.config["configurable"].get("checkpoint_id") == target_id:
continue # skip the checkpoint itself; we want its ancestors' writes
if tup.pending_writes:
for _, ch, value in tup.pending_writes:
if ch == channel:
result.append(value)
result.reverse() # list() yields newest→oldest; we want oldest→newest
return result
finally:
_DELTA_RECONSTRUCTION.active = False
async def aget_channel_writes(
self, config: RunnableConfig, channel: str
) -> List[Any]: # noqa: UP006
"""Async version of get_channel_writes."""
if getattr(_DELTA_RECONSTRUCTION, "active", False):
return []
_DELTA_RECONSTRUCTION.active = True
try:
result: list[Any] = []
target_id = config["configurable"].get("checkpoint_id")
async for tup in self.alist(config):
if tup.config["configurable"].get("checkpoint_id") == target_id:
continue
if tup.pending_writes:
for _, ch, value in tup.pending_writes:
if ch == channel:
result.append(value)
result.reverse()
return result
finally:
_DELTA_RECONSTRUCTION.active = False
def get_next_version(self, current: V | None, channel: None) -> V:
"""Generate the next version ID for a channel.
@@ -9,7 +9,7 @@ from collections import defaultdict
from collections.abc import AsyncIterator, Iterator, Sequence
from contextlib import AbstractAsyncContextManager, AbstractContextManager, ExitStack
from types import TracebackType
from typing import Any, cast
from typing import Any
from langchain_core.runnables import RunnableConfig
@@ -20,7 +20,8 @@ from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
DeltaChannelSentinel,
DeltaChainValue,
DeltaValue,
SerializerProtocol,
get_checkpoint_id,
get_checkpoint_metadata,
@@ -126,56 +127,68 @@ class InMemorySaver(
thread_id: str,
checkpoint_ns: str,
versions: ChannelVersions,
checkpoint_id: str = "",
) -> dict[str, Any]:
result: dict[str, Any] = {}
for k, ver in versions.items():
kk = (thread_id, checkpoint_ns, k, ver)
channel_values: dict[str, Any] = {}
delta_channels: list[str] = []
for k, v in versions.items():
kk = (thread_id, checkpoint_ns, k, v)
if kk not in self.blobs:
continue
vv = self.blobs[kk]
if vv[0] == "empty":
continue
result[k] = self.serde.loads_typed(vv)
return result
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 _resolve_delta_channels(
def _assemble_delta_chain(
self,
config: RunnableConfig,
channel_values: dict[str, Any],
) -> None:
"""Replace DeltaChannelSentinel entries with reconstructed write lists."""
for channel, value in list(channel_values.items()):
if isinstance(value, DeltaChannelSentinel):
channel_values[channel] = self.get_channel_writes(config, channel)
def get_channel_writes(self, config: RunnableConfig, channel: str) -> list[Any]:
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = config["configurable"].get("checkpoint_id", "")
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, {})
# Walk the parent chain newest→oldest collecting checkpoint IDs.
chain: list[str] = []
current: str | None = checkpoint_id
while current is not None:
entry = ns_storage.get(current)
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
chain.append(current)
_, _, parent = entry
current = parent
# Collect writes oldest→newest.
result: list[Any] = []
for cp_id in reversed(chain):
step_writes = self.writes.get((thread_id, checkpoint_ns, cp_id), {})
for (_task_id, _idx), (_, ch, serialized, _) in sorted(step_writes.items()):
if ch == channel:
result.append(self.serde.loads_typed(serialized))
return result
async def aget_channel_writes(
self, config: RunnableConfig, channel: str
) -> list[Any]:
return self.get_channel_writes(config, channel)
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.
@@ -198,17 +211,16 @@ class InMemorySaver(
checkpoint, metadata, parent_checkpoint_id = saved
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
channel_values = self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
)
self._resolve_delta_channels(config, channel_values)
return CheckpointTuple(
config=config,
checkpoint={
**checkpoint_,
"channel_values": channel_values,
"channel_values": self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
checkpoint_id,
),
},
metadata=self.serde.loads_typed(metadata),
pending_writes=[
@@ -232,27 +244,22 @@ class InMemorySaver(
checkpoint, metadata, parent_checkpoint_id = checkpoints[checkpoint_id]
writes = self.writes[(thread_id, checkpoint_ns, checkpoint_id)].values()
checkpoint_ = self.serde.loads_typed(checkpoint)
resolved_config = cast(
RunnableConfig,
{
return CheckpointTuple(
config={
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": checkpoint_ns,
"checkpoint_id": checkpoint_id,
}
},
)
channel_values = self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
)
self._resolve_delta_channels(resolved_config, channel_values)
return CheckpointTuple(
config=resolved_config,
checkpoint={
**checkpoint_,
"channel_values": channel_values,
"channel_values": self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
checkpoint_id,
),
},
metadata=self.serde.loads_typed(metadata),
pending_writes=[
@@ -347,28 +354,22 @@ class InMemorySaver(
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
list_config = cast(
RunnableConfig,
{
yield CheckpointTuple(
config={
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": checkpoint_ns,
"checkpoint_id": checkpoint_id,
}
},
)
channel_values = self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
)
self._resolve_delta_channels(list_config, channel_values)
yield CheckpointTuple(
config=list_config,
checkpoint={
**checkpoint_,
"channel_values": channel_values,
"channel_values": self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
checkpoint_id,
),
},
metadata=metadata,
parent_config=(
@@ -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"),
}
)
@@ -46,30 +46,11 @@ LC_REVIVER = Reviver()
EMPTY_BYTES = b""
logger = logging.getLogger(__name__)
# Dedup log warnings across process lifetime; cap bounds state if types are
# dynamically generated (also acts as a circuit breaker on warning volume).
# Dedup is best-effort: racing threads may each emit once for the same key,
# and warnings are silently dropped once _MAX_WARNED_TYPES is reached.
_MAX_WARNED_TYPES = 1000
_warned_unregistered_types: set[tuple[str, str]] = set()
_warned_blocked_types: set[tuple[str, str]] = set()
def _is_delta_value(obj: Any) -> bool:
from langgraph.checkpoint.base import DeltaValue # lazy import avoids circular dep
def _warn_once(
seen: set[tuple[str, str]], key: tuple[str, str], msg: str, *args: object
) -> None:
if key in seen or len(seen) >= _MAX_WARNED_TYPES:
return
seen.add(key)
logger.warning(msg, *args)
def _get_delta_sentinel_cls() -> type:
from langgraph.checkpoint.base import (
DeltaChannelSentinel,
) # lazy import avoids circular dep
return DeltaChannelSentinel
return isinstance(obj, DeltaValue)
class JsonPlusSerializer(SerializerProtocol):
@@ -264,8 +245,8 @@ class JsonPlusSerializer(SerializerProtocol):
return "bytes", obj
elif isinstance(obj, bytearray):
return "bytearray", obj
elif isinstance(obj, _get_delta_sentinel_cls()):
return "delta", b""
elif _is_delta_value(obj):
return "delta", _msgpack_enc({"d": obj.delta})
else:
try:
return "msgpack", _msgpack_enc(obj)
@@ -289,9 +270,12 @@ class JsonPlusSerializer(SerializerProtocol):
data_, ext_hook=self._unpack_ext_hook, option=ormsgpack.OPT_NON_STR_KEYS
)
elif type_ == "delta":
from langgraph.checkpoint.base import DeltaChannelSentinel
from langgraph.checkpoint.base import DeltaValue # lazy import
return DeltaChannelSentinel()
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:
@@ -565,9 +549,7 @@ def _create_msgpack_ext_hook(
"name": name,
}
)
_warn_once(
_warned_unregistered_types,
key,
logger.warning(
"Deserializing unregistered type %s.%s from checkpoint. "
"This will be blocked in a future version. "
"Set LANGGRAPH_STRICT_MSGPACK=true to block now, or add "
@@ -589,9 +571,7 @@ def _create_msgpack_ext_hook(
"name": name,
}
)
_warn_once(
_warned_blocked_types,
key,
logger.warning(
"Blocked deserialization of %s.%s - not in allowed_msgpack_modules. "
"Add to allowed_msgpack_modules to allow: [(%r, %r)]",
module,
-9
View File
@@ -29,8 +29,6 @@ from langgraph.checkpoint.serde.jsonplus import (
EXT_METHOD_SINGLE_ARG,
JsonPlusSerializer,
_msgpack_enc,
_warned_blocked_types,
_warned_unregistered_types,
)
@@ -104,13 +102,6 @@ def test_msgpack_method_pathlib_blocked_encrypted_strict(
class TestEncryptedSerializerMsgpackAllowlist:
"""Test msgpack allowlist behavior through EncryptedSerializer."""
@pytest.fixture(autouse=True)
def _reset_warned_types(self) -> None:
# Warning dedup state is process-global; reset per-test so each case
# sees a fresh slate and assertions about warning emission are stable.
_warned_unregistered_types.clear()
_warned_blocked_types.clear()
def test_safe_types_no_warning(self, caplog: pytest.LogCaptureFixture) -> None:
"""Test safe types deserialize without warnings through encryption."""
serde = _make_encrypted_serde()
+7 -20
View File
@@ -35,8 +35,6 @@ from langgraph.checkpoint.serde.jsonplus import (
JsonPlusSerializer,
_msgpack_enc,
_msgpack_ext_hook_to_json,
_warned_blocked_types,
_warned_unregistered_types,
)
from langgraph.store.base import Item
@@ -582,14 +580,6 @@ def test_msgpack_safe_types_no_warning(caplog: pytest.LogCaptureFixture) -> None
assert result is not None
@pytest.fixture(autouse=True)
def _reset_warned_types() -> None:
# Warning dedup state is process-global; reset per-test so each case sees
# a fresh slate and assertions about warning emission are stable.
_warned_unregistered_types.clear()
_warned_blocked_types.clear()
def test_msgpack_pydantic_warns_by_default(caplog: pytest.LogCaptureFixture) -> None:
"""Pydantic models not in allowlist should log warning but still deserialize."""
current = _lg_msgpack.STRICT_MSGPACK_ENABLED
@@ -605,12 +595,6 @@ def test_msgpack_pydantic_warns_by_default(caplog: pytest.LogCaptureFixture) ->
assert "unregistered type" in caplog.text.lower()
assert "allowed_msgpack_modules" in caplog.text
assert result == obj
# Second deserialization of the same type should NOT produce another warning
caplog.clear()
result2 = serde.loads_typed(dumped)
assert "unregistered type" not in caplog.text.lower()
assert result2 == obj
_lg_msgpack.STRICT_MSGPACK_ENABLED = current
@@ -655,6 +639,7 @@ def test_msgpack_allowlist_silences_warning(caplog: pytest.LogCaptureFixture) ->
def test_msgpack_none_blocks_unregistered(caplog: pytest.LogCaptureFixture) -> None:
"""allowed_msgpack_modules=None should block unregistered types."""
serde = JsonPlusSerializer(allowed_msgpack_modules=None)
obj = MyPydantic(foo="test", bar=42, inner=InnerPydantic(hello="world"))
@@ -672,6 +657,7 @@ def test_msgpack_allowlist_blocks_non_listed(
caplog: pytest.LogCaptureFixture,
) -> None:
"""Allowlists should block unregistered types even if msgpack is enabled."""
serde = JsonPlusSerializer(
allowed_msgpack_modules=[("tests.test_jsonplus", "MyPydantic")]
)
@@ -999,13 +985,14 @@ def test_msgpack_nested_pydantic_serializes_as_dict(
assert result == obj
def test_delta_channel_sentinel_serde_round_trip() -> None:
from langgraph.checkpoint.base import DeltaChannelSentinel
def test_delta_value_serde_round_trip() -> None:
from langgraph.checkpoint.base import DeltaValue
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
serde = JsonPlusSerializer()
original = DeltaChannelSentinel()
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, DeltaChannelSentinel)
assert isinstance(loaded, DeltaValue)
assert loaded.delta == original.delta
+35 -50
View File
@@ -12,25 +12,13 @@ from langgraph.checkpoint.base import (
empty_checkpoint,
)
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.checkpoint.serde.jsonplus import (
JsonPlusSerializer,
_warned_blocked_types,
_warned_unregistered_types,
)
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
class MemoryPydantic(BaseModel):
foo: str
@pytest.fixture(autouse=True)
def _reset_warned_types() -> None:
# Warning dedup state is process-global; reset per-test so each case sees
# a fresh slate and assertions about warning emission are stable.
_warned_unregistered_types.clear()
_warned_blocked_types.clear()
class TestMemorySaver:
@pytest.fixture(autouse=True)
def setup(self) -> None:
@@ -323,10 +311,11 @@ def test_memory_saver_with_allowlist_proxy_isolated() -> None:
class TestInMemorySaverDeltaChannel:
def test_load_blobs_returns_sentinel_for_delta_channel(self) -> None:
"""_load_blobs returns DeltaChannelSentinel for delta channels (reconstruction deferred)."""
def test_load_blobs_assembles_delta_chain(self) -> None:
"""_load_blobs returns DeltaChainValue for delta channels, not raw DeltaValue."""
from langgraph.checkpoint.base import (
DeltaChannelSentinel,
DeltaChainValue,
DeltaValue,
empty_checkpoint,
)
@@ -335,59 +324,55 @@ class TestInMemorySaverDeltaChannel:
thread_id, ns, channel = "t1", "", "messages"
v1 = "00000000000000000000000000000001.0000000000000000"
v2 = "00000000000000000000000000000002.0000000000000000"
sentinel = DeltaChannelSentinel()
saver.blobs[(thread_id, ns, channel, v1)] = serde.dumps_typed(sentinel)
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: v1})
result = saver._load_blobs(thread_id, ns, {channel: v2}, "cp2")
assert channel in result
assert isinstance(result[channel], DeltaChannelSentinel)
chain = result[channel]
assert isinstance(chain, DeltaChainValue)
assert chain.deltas == [[{"content": "hi"}], [{"content": "bye"}]]
def test_get_channel_writes_collects_writes(self) -> None:
"""get_channel_writes collects per-step writes oldest→newest."""
from langgraph.checkpoint.base import empty_checkpoint
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"
cp2 = empty_checkpoint()
cp2["id"] = "cp2"
cp1["channel_versions"][channel] = v1
saver.storage[thread_id][ns] = {
"cp1": (serde.dumps_typed(cp1), serde.dumps_typed({}), None),
"cp2": (serde.dumps_typed(cp2), serde.dumps_typed({}), "cp1"),
"cp1": (serde.dumps_typed(cp1), serde.dumps_typed({}), None)
}
# cp1 has a write for channel
saver.writes[(thread_id, ns, "cp1")][("task1", 0)] = (
"task1",
channel,
serde.dumps_typed({"content": "hi"}),
"",
)
# cp2 has a write for channel
saver.writes[(thread_id, ns, "cp2")][("task2", 0)] = (
"task2",
channel,
serde.dumps_typed({"content": "bye"}),
"",
)
config: RunnableConfig = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": ns,
"checkpoint_id": "cp2",
}
}
result = saver.get_channel_writes(config, channel)
assert result == [{"content": "hi"}, {"content": "bye"}]
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"}]]
@@ -5,5 +5,5 @@ description = "Test for prerelease stuff"
readme = "README.md"
requires-python = ">=3.10"
dependencies = [
"langchain-openai==1.1.14"
"langchain-openai==1.0.1"
]
@@ -5,7 +5,7 @@ description = "Test for prerelease stuff"
readme = "README.md"
requires-python = ">=3.10"
dependencies = [
"langchain-openai==1.1.14",
"langchain-openai==1.0.0a2",
"langchain-anthropic==1.0.0a5",
"langgraph==1.1.5"
]
@@ -5,7 +5,7 @@ description = "Test for prerelease stuff"
readme = "README.md"
requires-python = ">=3.10"
dependencies = [
"langchain-openai==1.1.14",
"langchain-openai==1.0.0a2",
"langgraph==1.1.2",
"langchain_community>=0.3.0",
]
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.23"
__version__ = "0.4.22"
+1 -1
View File
@@ -23,7 +23,7 @@ dependencies = [
path = "langgraph_cli/__init__.py"
[project.optional-dependencies]
inmem = [
"langgraph-api>=0.5.35,<0.9.0 ; python_version >= '3.11'",
"langgraph-api>=0.5.35,<0.8.0 ; python_version >= '3.11'",
"langgraph-runtime-inmem>=0.7 ; python_version >= '3.11'",
]
+383 -465
View File
File diff suppressed because it is too large Load Diff
+18
View File
@@ -245,6 +245,15 @@ class _GraphCallbackManager(BaseCallbackManager):
run_id=run_id,
)
def add_handler(
self,
handler: BaseCallbackHandler,
inherit: bool = True, # noqa: FBT001,FBT002
) -> None:
if not isinstance(handler, GraphCallbackHandler):
raise TypeError("handlers must inherit GraphCallbackHandler")
super().add_handler(handler, inherit=inherit)
def copy(
self,
*,
@@ -312,6 +321,15 @@ class _AsyncGraphCallbackManager(BaseCallbackManager):
run_id=run_id,
)
def add_handler(
self,
handler: BaseCallbackHandler,
inherit: bool = True, # noqa: FBT001,FBT002
) -> None:
if not isinstance(handler, GraphCallbackHandler):
raise TypeError("handlers must inherit GraphCallbackHandler")
super().add_handler(handler, inherit=inherit)
def copy(
self,
*,
+90 -31
View File
@@ -3,7 +3,7 @@ from __future__ import annotations
from collections.abc import Callable, Sequence
from typing import Any, Generic
from langgraph.checkpoint.base import DeltaChannelSentinel
from langgraph.checkpoint.base import DeltaChainValue, DeltaValue
from typing_extensions import Self
from langgraph._internal._typing import MISSING
@@ -14,41 +14,64 @@ from langgraph.errors import EmptyChannelError
__all__ = ("DeltaChannel",)
class DeltaChannel(
Generic[Value], BaseChannel[list[Value], Value, DeltaChannelSentinel]
):
"""A channel that stores only a sentinel in checkpoints; per-step writes are
stored in checkpoint_writes and replayed through the operator at load time.
class DeltaChannel(Generic[Value], BaseChannel[list[Value], Value, DeltaValue]):
"""A channel that stores only per-step write deltas in checkpoints.
Use with append-style reducers (e.g. `add_messages`) on long-running threads
to eliminate O(N²) blob growth — storage is O(N) using the writes table that
every checkpointer already maintains.
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 dedicated implementations
(InMemorySaver, PostgresSaver) reconstruct in one pass; others fall back to
walking the checkpoint list.
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")
__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>"
@@ -65,14 +88,18 @@ class DeltaChannel(
return self.typ | list[self.typ] # type: ignore[name-defined]
def copy(self) -> Self:
new = DeltaChannel(self.operator)
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)
new = DeltaChannel(self.operator, snapshot_every=self.snapshot_every)
new.typ = self.typ
new.key = self.key
if checkpoint is MISSING:
@@ -80,21 +107,29 @@ class DeltaChannel(
new.value = new.typ()
except Exception:
new.value = []
elif isinstance(checkpoint, list):
# Flat list of write values (oldest→newest) from get_channel_writes.
try:
value: Any = new.typ()
except Exception:
value = []
for write in checkpoint:
value = new.operator(value, write)
new.value = 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:
# Backward compat: plain accumulated value (e.g. from a migrated thread).
try:
new.value = list(checkpoint)
except Exception:
new.value = []
# 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:
@@ -119,10 +154,13 @@ class DeltaChannel(
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]:
@@ -133,5 +171,26 @@ class DeltaChannel(
def is_available(self) -> bool:
return self.value is not MISSING
def checkpoint(self) -> DeltaChannelSentinel:
return DeltaChannelSentinel()
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
+1 -11
View File
@@ -831,18 +831,8 @@ class PregelLoop:
# parent. For forks (source=update/fork), use the fork's parent
# checkpoint ID since the fork was created after the subgraph's
# checkpoints from the original execution.
#
# Only gate on is_time_traveling (not is_replaying). When the
# client resumes with an explicit checkpoint_id that happens to
# point at the current head (e.g. LangGraph Studio sending
# `checkpoint: {checkpoint_id}` alongside Command(resume=...)),
# is_replaying is True but is_time_traveling is False. In that
# case subgraphs should load their latest checkpoint normally,
# not go through ReplayState's before-bound lookup which would
# miss subgraph checkpoints created during processing of the
# current parent step.
replay_state: ReplayState | None = None
if is_time_traveling:
if self.is_replaying:
replay_checkpoint_id = self.checkpoint["id"]
if (
self.checkpoint_metadata.get("source")
+14 -6
View File
@@ -14,7 +14,7 @@ from langchain_core.messages import BaseMessage
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, LLMResult
from pydantic import BaseModel
from langgraph._internal._constants import NS_SEP
from langgraph._internal._constants import NS_END, NS_SEP
from langgraph.constants import TAG_HIDDEN, TAG_NOSTREAM
from langgraph.pregel.protocol import StreamChunk
from langgraph.types import Command
@@ -132,15 +132,23 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
**kwargs: Any,
) -> Any:
if metadata and (not tags or (TAG_NOSTREAM not in tags)):
ns = tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP))[
:-1
]
task_checkpoint_ns = cast(str, metadata["langgraph_checkpoint_ns"])
checkpoint_ns = (
f"{task_checkpoint_ns.rsplit(NS_END, 1)[0]}{NS_END}"
if NS_END in task_checkpoint_ns
else task_checkpoint_ns
)
ns = tuple(task_checkpoint_ns.split(NS_SEP))[:-1]
if not self.subgraphs and len(ns) > 0 and ns != self.parent_ns:
return
stream_metadata = dict(metadata)
stream_metadata["langgraph_checkpoint_ns"] = checkpoint_ns
# Preserve backwards-compatible streamed checkpoint metadata shape.
stream_metadata["checkpoint_ns"] = checkpoint_ns
if tags:
if filtered_tags := [t for t in tags if not t.startswith("seq:step")]:
metadata["tags"] = filtered_tags
self.metadata[run_id] = (ns, metadata)
stream_metadata["tags"] = filtered_tags
self.metadata[run_id] = (ns, stream_metadata)
def on_llm_new_token(
self,
+30 -60
View File
@@ -1,11 +1,9 @@
from __future__ import annotations
import ast
import functools
import inspect
import re
import textwrap
import types
from collections.abc import Callable
from typing import Any
@@ -66,74 +64,46 @@ def find_subgraph_pregel(candidate: Runnable) -> PregelProtocol | None:
return None
@functools.lru_cache(maxsize=256)
def _get_nonlocal_names(code: types.CodeType) -> frozenset[str]:
"""Return the set of nonlocal variable names referenced by a function.
Cached by code object so the expensive source fetch + AST parse only
happens once per unique function definition across repeated graph compiles.
Args:
code: The code object of the function to analyse.
Returns:
Frozenset of variable names that the function reads from its enclosing
scope (free variables and globals referenced in function bodies).
"""
try:
source = inspect.getsource(code)
tree = ast.parse(textwrap.dedent(source))
visitor = FunctionNonLocals()
visitor.visit(tree)
return frozenset(visitor.nonlocals)
except (SyntaxError, TypeError, OSError, SystemError):
return frozenset()
def get_function_nonlocals(func: Callable) -> list[Any]:
"""Get the nonlocal variables accessed by a function.
The expensive source-parsing step is cached by code object; only the
cheap closure-variable lookup runs on every call.
Args:
func: The function to check.
Returns:
List[Any]: The nonlocal variables accessed by the function.
"""
actual_func = (
func.__wrapped__
if hasattr(func, "__wrapped__") and callable(func.__wrapped__)
else func
)
# Fast path: no free variables means nothing to scan.
if not actual_func.__code__.co_freevars:
return []
nonlocal_names = _get_nonlocal_names(actual_func.__code__)
if not nonlocal_names:
return []
closure = inspect.getclosurevars(actual_func)
candidates = {**closure.globals, **closure.nonlocals}
values: list[Any] = []
for k, v in candidates.items():
if k in nonlocal_names:
values.append(v)
for kk in nonlocal_names:
if "." in kk and kk.startswith(k):
vv = v
for part in kk.split(".")[1:]:
if vv is None:
break
else:
try:
vv = getattr(vv, part)
except AttributeError:
try:
code = inspect.getsource(func)
tree = ast.parse(textwrap.dedent(code))
visitor = FunctionNonLocals()
visitor.visit(tree)
values: list[Any] = []
closure = (
inspect.getclosurevars(func.__wrapped__)
if hasattr(func, "__wrapped__") and callable(func.__wrapped__)
else inspect.getclosurevars(func)
)
candidates = {**closure.globals, **closure.nonlocals}
for k, v in candidates.items():
if k in visitor.nonlocals:
values.append(v)
for kk in visitor.nonlocals:
if "." in kk and kk.startswith(k):
vv = v
for part in kk.split(".")[1:]:
if vv is None:
break
else:
values.append(vv)
else:
try:
vv = getattr(vv, part)
except AttributeError:
break
else:
values.append(vv)
except (SyntaxError, TypeError, OSError, SystemError):
return []
return values
+2 -2
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph"
version = "1.1.9"
version = "1.1.7a2"
description = "Building stateful, multi-actor applications with LLMs"
authors = []
requires-python = ">=3.10"
@@ -24,7 +24,7 @@ classifiers = [
'Programming Language :: Python :: 3.13',
]
dependencies = [
"langchain-core>=1.3.0,<2",
"langchain-core==1.3.0a2",
"langgraph-checkpoint>=2.1.0,<5.0.0",
"langgraph-sdk>=0.3.0,<0.4.0",
"langgraph-prebuilt>=1.0.9,<1.1.0",
+189 -66
View File
@@ -121,22 +121,26 @@ def test_untracked_value() -> None:
def test_delta_channel_basic_two_steps() -> None:
from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.base import DeltaChannelSentinel
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, DeltaChannelSentinel)
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 isinstance(d2, DeltaChannelSentinel)
assert len(d2.delta) == 1
ch.after_checkpoint("v2")
# Full accumulated value is preserved in memory
assert len(ch.get()) == 2
@@ -144,21 +148,40 @@ def test_delta_channel_basic_two_steps() -> None:
assert ch.get()[1].content == "hello"
def test_delta_channel_from_checkpoint_writes_list() -> None:
"""from_checkpoint with a flat list of individual writes replays them through the operator."""
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)
# Each element is one write value (as stored in checkpoint_writes)
writes = [
HumanMessage(content="hi", id="h1"),
AIMessage(content="hello", id="a1"),
HumanMessage(content="bye", id="h2"),
]
ch = spec.from_checkpoint(writes)
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"
@@ -172,45 +195,68 @@ def test_delta_channel_from_checkpoint_backwards_compat() -> None:
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
# Old BinaryOperatorAggregate checkpoint: plain list treated as backward compat
# 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() -> None:
def test_delta_channel_overwrite_resets_chain() -> None:
from langchain_core.messages import HumanMessage
from langgraph.checkpoint.base import DeltaChannelSentinel
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, DeltaChannelSentinel)
# After overwrite, value is reset to only the new message
assert len(ch.get()) == 1
assert ch.get()[0].content == "new"
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_remove_message_and_replay() -> None:
"""RemoveMessage must round-trip correctly when writes are replayed."""
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"),
@@ -218,63 +264,127 @@ def test_delta_channel_remove_message_and_replay() -> None:
# 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 writes list from scratch — must reproduce the post-remove state
writes = [
HumanMessage(content="hi", id="h1"),
AIMessage(content="hello", id="a1"),
RemoveMessage(id="a1"),
]
ch2 = spec.from_checkpoint(writes)
# 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_and_replay() -> None:
"""Updating a message by ID must round-trip correctly through writes replay."""
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 writes — must produce the updated message, not the original
writes = [
HumanMessage(content="original", id="h1"),
HumanMessage(content="updated", id="h1"),
]
ch2 = spec.from_checkpoint(writes)
# 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_checkpoint_returns_sentinel() -> None:
"""checkpoint() always returns DeltaChannelSentinel regardless of state."""
from langgraph.checkpoint.base import DeltaChannelSentinel
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
ch = DeltaChannel(add_messages).from_checkpoint(MISSING)
assert isinstance(ch.checkpoint(), DeltaChannelSentinel)
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")
from langchain_core.messages import HumanMessage
# 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}")
ch.update([HumanMessage(content="hi", id="h1")])
assert isinstance(ch.checkpoint(), DeltaChannelSentinel)
# 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_inmemory_saver_assembles_writes() -> None:
"""InMemorySaver assembles writes from checkpoint_writes inside get_tuple."""
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
@@ -304,16 +414,14 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None:
graph.invoke({"messages": [HumanMessage(content="hi", id="h1")]}, config)
graph.invoke({"messages": [HumanMessage(content="bye", id="h2")]}, config)
# get_tuple must return a resolved list (not DeltaChannelSentinel)
from langgraph.checkpoint.base import DeltaChannelSentinel
# 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"], DeltaChannelSentinel
)
assert isinstance(saved.checkpoint["channel_values"]["messages"], list)
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
@@ -343,39 +451,48 @@ def test_delta_channel_dict_reducer_fresh_channel() -> None:
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 DeltaChannelSentinel
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, DeltaChannelSentinel)
assert isinstance(d1, DeltaValue)
assert d1.delta == [{"a": 1}]
ch.after_checkpoint("v1", checkpoint_id="cid1")
ch.update([{"b": 2}])
d2 = ch.checkpoint()
assert isinstance(d2, DeltaChannelSentinel)
assert d2.delta == [{"b": 2}]
ch.after_checkpoint("v2")
assert ch.get() == {"a": 1, "b": 2}
def test_delta_channel_dict_reducer_writes_reconstruction() -> None:
"""from_checkpoint with a writes list replays correctly through a dict merge reducer."""
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)
# Each element is one write value (oldest→newest)
writes = [{"a": 1}, {"b": 2}, {"c": 3}]
ch = spec.from_checkpoint(writes)
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:
@@ -389,19 +506,25 @@ def test_delta_channel_dict_reducer_with_deletions() -> None:
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 writes reconstruction produces the same result
writes = [
{"file1.py": "content1", "file2.py": "content2"},
{"file1.py": None, "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(writes)
ch2 = spec.from_checkpoint(chain)
assert ch2.get() == {"file2.py": "content2", "file3.py": "content3"}
@@ -8,11 +8,6 @@ Simulates realistic multi-turn conversations with paragraph-length messages
Token estimates: 1 token ≈ 4 chars; each turn ≈ 200 tokens (human + AI).
A 1M-token conversation ≈ 5,000 turns of realistic messages.
DeltaChannel stores only a zero-byte sentinel in checkpoint_blobs; the actual
write data lives in checkpoint_writes (already stored there). Reconstruction
walks the parent chain and replays writes through the operator — O(N) total
storage vs O(N²) for plain add_messages.
"""
from __future__ import annotations
@@ -47,6 +42,8 @@ try:
except ImportError:
_POSTGRES_AVAILABLE = False
SNAPSHOT_EVERY = 50
# ---------------------------------------------------------------------------
# Realistic message payload (~100 tokens / ~400 chars each)
# ---------------------------------------------------------------------------
@@ -124,6 +121,10 @@ class DeltaState(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages)]
class DeltaSnapshotState(TypedDict):
messages: Annotated[list, DeltaChannel(add_messages, snapshot_every=SNAPSHOT_EVERY)]
# ---------------------------------------------------------------------------
# Graph factory
# ---------------------------------------------------------------------------
@@ -219,7 +220,7 @@ def _approx_tokens(n_turns: int) -> str:
# 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, 500]
TURN_COUNTS = [10, 25, 50, 100]
def _checkpointer_factories() -> list[tuple[str, Any]]:
@@ -238,13 +239,7 @@ def run_benchmark() -> None:
checkpointers: list[tuple[str, Any]] = [("InMemory", None)]
if _POSTGRES_AVAILABLE:
try:
import psycopg
psycopg.connect(_POSTGRES_URI).close()
checkpointers.append(("Postgres (recursive CTE)", "postgres"))
except Exception:
pass
checkpointers.append(("Postgres (recursive CTE)", "postgres"))
for cp_label, cp_hint in checkpointers:
print(f"--- Checkpointer: {cp_label} ---")
@@ -282,28 +277,30 @@ def _run_benchmark_for_checkpointer(cp_hint: Any) -> None:
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)
rows.append((turns, b_bytes, d_bytes, b_rt, d_rt))
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 = 70
W = 80
print("Storage (checkpoint blob bytes)")
print("=" * W)
print(
f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12} {'savings':>8}"
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, b_rt, d_rt in rows:
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':>8}"
f"{turns:>6} {_approx_tokens(turns):>10} {'n/a':>12} {'n/a':>12} {'n/a':>12} {'n/a':>8}"
)
else:
ratio = b_bytes / d_bytes if d_bytes else float("inf")
storage_results.append((turns, b_bytes, d_bytes, ratio))
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} "
f"{_fmt_bytes(b_bytes):>12} {_fmt_bytes(d_bytes):>12} {_fmt_bytes(s_bytes):>12} "
f"{ratio:>7.0f}x"
)
print("=" * W)
@@ -312,30 +309,35 @@ def _run_benchmark_for_checkpointer(cp_hint: Any) -> None:
# ── 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}")
print(
f"{'turns':>6} {'ctx size':>10} {'add_msgs':>12} {'delta':>12} {'delta+snap':>12}"
)
print("-" * W)
for turns, b_bytes, d_bytes, b_rt, d_rt in rows:
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"
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:
turns, b_bytes, d_bytes, ratio = storage_results[-1]
b_rt = rows[-1][-2]
d_rt = rows[-1][-1]
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(d_bytes)} ({ratio:.0f}x less storage); "
f"read {b_rt * 1000:.1f}ms → {d_rt * 1000:.1f}ms"
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(" add_msgs = Annotated[list, add_messages] — O(N²) storage")
print(
" delta = DeltaChannel(add_messages) — O(N) storage, full chain replay"
" 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()
@@ -357,10 +359,15 @@ def test_delta_channel_benchmark(capsys: Any) -> None:
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}"
)
# ---------------------------------------------------------------------------
@@ -275,70 +275,3 @@ def test_graph_callbacks_accept_base_callback_manager() -> None:
assert "__interrupt__" in first
assert len(graph_handler.interrupt_events) == 1
def test_non_graph_handler_via_add_handler_does_not_crash() -> None:
"""Non-GraphCallbackHandler added via add_handler should not raise.
Libraries like opentelemetry-instrumentation-langchain monkey-patch
BaseCallbackManager.__init__ and inject handlers via add_handler().
These handlers inherit from BaseCallbackHandler, not
GraphCallbackHandler. They must be silently accepted — graph lifecycle
events will simply not be dispatched to them.
"""
from langgraph.callbacks import _GraphCallbackManager
manager = _GraphCallbackManager()
plain_handler = _LangChainCustomEventHandler()
manager.add_handler(plain_handler, inherit=True)
assert plain_handler in manager.handlers
def test_non_graph_handler_does_not_receive_lifecycle_events() -> None:
"""Non-GraphCallbackHandler added alongside a GraphCallbackHandler
should not interfere with lifecycle event dispatch."""
graph = _build_interrupt_graph()
graph_handler = _GraphEventHandler()
plain_handler = _LangChainCustomEventHandler()
config = {
"configurable": {"thread_id": "graph-callback-mixed-handlers"},
"callbacks": [plain_handler, graph_handler],
}
first = graph.invoke({"answer": None}, config)
assert "__interrupt__" in first
assert len(graph_handler.interrupt_events) == 1
assert plain_handler.events == []
resumed = graph.invoke(Command(resume="done"), config)
assert resumed == {"answer": "done"}
assert len(graph_handler.resume_events) == 1
assert plain_handler.events == []
@pytest.mark.anyio
@NEEDS_CONTEXTVARS
async def test_non_graph_handler_does_not_receive_lifecycle_events_async() -> None:
"""Async variant: non-GraphCallbackHandler should not interfere."""
graph = _build_interrupt_graph()
graph_handler = _GraphEventHandler()
plain_handler = _LangChainCustomEventHandler()
config = {
"configurable": {"thread_id": "graph-callback-mixed-handlers-async"},
"callbacks": [plain_handler, graph_handler],
}
first = await graph.ainvoke({"answer": None}, config)
assert "__interrupt__" in first
assert len(graph_handler.interrupt_events) == 1
assert plain_handler.events == []
resumed = await graph.ainvoke(Command(resume="done"), config)
assert resumed == {"answer": "done"}
assert len(graph_handler.resume_events) == 1
assert plain_handler.events == []
@@ -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)
-66
View File
@@ -1113,72 +1113,6 @@ def test_subgraph_interrupt_replay_from_parent_then_resume(
]
def test_subgraph_interrupt_resume_with_explicit_head_checkpoint_id(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Resume with Command(resume=...) plus the current head checkpoint_id
in config. The subgraph must continue from the interrupted node, not
restart from scratch. Explicit checkpoint_id triggers is_replaying but
this is a resume, not a time-travel, so ReplayState should not apply."""
called: list[str] = []
def step_a(state: State) -> State:
called.append("step_a")
return {"value": ["sub_a"]}
def ask_human(state: State) -> State:
called.append("ask_human")
answer = interrupt("Provide input:")
return {"value": [f"human:{answer}"]}
def step_b(state: State) -> State:
called.append("step_b")
return {"value": ["sub_b"]}
subgraph = (
StateGraph(State)
.add_node("step_a", step_a)
.add_node("ask_human", ask_human)
.add_node("step_b", step_b)
.add_edge(START, "step_a")
.add_edge("step_a", "ask_human")
.add_edge("ask_human", "step_b")
.compile(checkpointer=True)
)
graph = (
StateGraph(State)
.add_node("subgraph_node", subgraph)
.add_edge(START, "subgraph_node")
.compile(checkpointer=sync_checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# Run until interrupt fires in subgraph
graph.invoke({"value": []}, config)
assert called == ["step_a", "ask_human"]
# Resume with explicit head checkpoint_id in config
head_checkpoint_id = graph.get_state(config).config["configurable"][
"checkpoint_id"
]
called.clear()
resume_config = {
"configurable": {
"thread_id": "1",
"checkpoint_id": head_checkpoint_id,
"checkpoint_ns": "",
}
}
result = graph.invoke(Command(resume="answer"), resume_config)
assert called == ["ask_human", "step_b"]
assert "__interrupt__" not in result
assert result["value"] == ["sub_a", "human:answer", "sub_b"]
def test_subgraph_replay_loads_accumulated_state_then_resume(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
-46
View File
@@ -427,49 +427,3 @@ def test_callback_manager_copies_configurable_ids_to_tracing_metadata() -> None:
"thread_id": "th-123",
"user_id": "uid-1",
}
def test_get_nonlocal_names_cached_by_code_object() -> None:
"""_get_nonlocal_names caches by code object so repeated calls are cheap."""
from langgraph.pregel._utils import _get_nonlocal_names
x = 1
def my_func() -> int:
return x
result1 = _get_nonlocal_names(my_func.__code__)
result2 = _get_nonlocal_names(my_func.__code__)
# Same frozenset instance returned (cache hit)
assert result1 is result2
assert "x" in result1
def test_get_function_nonlocals_fast_path_no_freevars() -> None:
"""Functions with no free variables return [] without AST parsing."""
from langgraph.pregel._utils import _get_nonlocal_names, get_function_nonlocals
cache_info_before = _get_nonlocal_names.cache_info()
def pure_func(a: int, b: int) -> int:
return a + b
result = get_function_nonlocals(pure_func)
# Should have returned early without touching the cache
assert result == []
assert _get_nonlocal_names.cache_info().misses == cache_info_before.misses
def test_get_function_nonlocals_returns_closure_values() -> None:
"""get_function_nonlocals correctly extracts values from closures."""
from langgraph.pregel._utils import get_function_nonlocals
sentinel = object()
def my_func() -> object:
return sentinel
result = get_function_nonlocals(my_func)
assert sentinel in result
+7 -7
View File
@@ -1348,7 +1348,7 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.0"
version = "1.3.0a2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
@@ -1360,14 +1360,14 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/92/fe/20190232d9b513242899dbb0c2bb77e31b4d61e343743adbe90ebc2603d2/langchain_core-1.3.0.tar.gz", hash = "sha256:14a39f528bf459aa3aa40d0a7f7f1bae7520d435ef991ae14a4ceb74d8c49046", size = 860755, upload-time = "2026-04-17T14:51:38.298Z" }
sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/f8/e2/dbfa347aa072a6dc4cd38d6f9ebfc730b4c14c258c47f480f4c5c546f177/langchain_core-1.3.0-py3-none-any.whl", hash = "sha256:baf16ee028475df177b9ab8869a751c79406d64a6f12125b93802991b566cced", size = 515140, upload-time = "2026-04-17T14:51:36.274Z" },
{ url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" },
]
[[package]]
name = "langgraph"
version = "1.1.9"
version = "1.1.7a2"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -1439,7 +1439,7 @@ test = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = ">=1.3.0,<2" },
{ name = "langchain-core", specifier = "==1.3.0a2" },
{ name = "langgraph-checkpoint", editable = "../checkpoint" },
{ name = "langgraph-prebuilt", editable = "../prebuilt" },
{ name = "langgraph-sdk", editable = "../sdk-py" },
@@ -1706,7 +1706,7 @@ inmem = [
requires-dist = [
{ name = "click", specifier = ">=8.1.7" },
{ name = "httpx", specifier = ">=0.24.0" },
{ name = "langgraph-api", marker = "python_full_version >= '3.11' and extra == 'inmem'", specifier = ">=0.5.35,<0.9.0" },
{ name = "langgraph-api", marker = "python_full_version >= '3.11' and extra == 'inmem'", specifier = ">=0.5.35,<0.8.0" },
{ name = "langgraph-runtime-inmem", marker = "python_full_version >= '3.11' and extra == 'inmem'", specifier = ">=0.7" },
{ name = "langgraph-sdk", marker = "python_full_version >= '3.11'", specifier = ">=0.1.0" },
{ name = "pathspec", specifier = ">=0.11.0" },
@@ -1742,7 +1742,7 @@ test = [
[[package]]
name = "langgraph-prebuilt"
version = "1.0.10"
version = "1.0.9"
source = { editable = "../prebuilt" }
dependencies = [
{ name = "langchain-core" },
+1 -20
View File
@@ -808,7 +808,6 @@ class ToolNode(RunnableCallable):
context=runtime.context,
store=runtime.store,
stream_writer=runtime.stream_writer,
tools=list(self.tools_by_name.values()),
execution_info=runtime.execution_info,
server_info=runtime.server_info,
)
@@ -843,7 +842,6 @@ class ToolNode(RunnableCallable):
context=runtime.context,
store=runtime.store,
stream_writer=runtime.stream_writer,
tools=list(self.tools_by_name.values()),
execution_info=runtime.execution_info,
server_info=runtime.server_info,
)
@@ -1578,7 +1576,6 @@ class ToolRuntime(_DirectlyInjectedToolArg, Generic[ContextT, StateT]):
- `context`: Runtime context (shared with `Runtime`)
- `store`: `BaseStore` instance for persistent storage (shared with `Runtime`)
- `stream_writer`: `StreamWriter` for streaming output (shared with `Runtime`)
- `tools`: List of all available `BaseTool` instances
No `Annotated` wrapper is needed - just use `runtime: ToolRuntime`
as a parameter.
@@ -1621,7 +1618,6 @@ class ToolRuntime(_DirectlyInjectedToolArg, Generic[ContextT, StateT]):
context: ContextT
config: RunnableConfig
stream_writer: StreamWriter
tools: list[BaseTool]
tool_call_id: str | None
store: BaseStore | None
execution_info: ExecutionInfo | None = None
@@ -1842,17 +1838,9 @@ def _get_injection_from_type(
return None
# Cache keyed by tool object identity. Stores (tool, result) to keep a strong
# reference that prevents GC from reusing the id for a different object.
_INJECTED_ARGS_CACHE: dict[int, tuple[BaseTool, _InjectedArgs]] = {}
def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs:
"""Extract all injected arguments from tool in a single pass.
Results are cached by tool identity so the expensive type-hint and schema
inspection only runs once per unique tool object across ToolNode instances.
This function analyzes both the tool's input schema and function signature
to identify all arguments that should be injected (state, store, runtime).
@@ -1862,11 +1850,6 @@ def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs:
Returns:
_InjectedArgs structure containing all detected injections.
"""
tool_id = id(tool)
entry = _INJECTED_ARGS_CACHE.get(tool_id)
if entry is not None and entry[0] is tool:
return entry[1]
# Get annotations from both schema and function signature
full_schema = tool.get_input_schema()
schema_annotations = get_all_basemodel_annotations(full_schema)
@@ -1912,12 +1895,10 @@ def _get_all_injected_args(tool: BaseTool) -> _InjectedArgs:
if _get_injection_from_type(type_, ToolRuntime):
runtime_arg = name
result = _InjectedArgs(
return _InjectedArgs(
state=state_args,
store=store_arg,
runtime=runtime_arg,
all_injected_keys=all_injected_keys,
_optional_state_args=_optional_state_args,
)
_INJECTED_ARGS_CACHE[tool_id] = (tool, result)
return result
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-prebuilt"
version = "1.0.10"
version = "1.0.9"
description = "Library with high-level APIs for creating and executing LangGraph agents and tools."
authors = []
requires-python = ">=3.10"
@@ -69,7 +69,6 @@ def _create_config_with_runtime(store=None, state=None):
context={},
store=store,
stream_writer=None,
tools=[],
tool_call_id="test_id",
)
return {
+8 -29
View File
@@ -2016,8 +2016,8 @@ async def test_tool_node_inject_runtime_dynamic_tool_via_wrap_tool_call_async()
assert tool_message.tool_call_id == "call_dynamic_2"
def test_tool_runtime_forwards_execution_info_server_info_and_tools() -> None:
"""Test that execution_info, server_info, and tools are forwarded from Runtime to ToolRuntime."""
def test_tool_runtime_forwards_execution_info_and_server_info() -> None:
"""Test that execution_info and server_info are forwarded from Runtime to ToolRuntime."""
from langgraph.runtime import ExecutionInfo, ServerInfo
exec_info = ExecutionInfo(
@@ -2043,15 +2043,9 @@ def test_tool_runtime_forwards_execution_info_server_info_and_tools() -> None:
"""Tool that captures runtime info."""
captured["execution_info"] = runtime.execution_info
captured["server_info"] = runtime.server_info
captured["tools"] = runtime.tools
return "ok"
@dec_tool
def other_tool(y: int) -> str:
"""Another tool available to the runtime."""
return str(y)
node = ToolNode([info_tool, other_tool])
node = ToolNode([info_tool])
tool_call = {
"name": "info_tool",
"args": {"x": 1},
@@ -2060,21 +2054,17 @@ def test_tool_runtime_forwards_execution_info_server_info_and_tools() -> None:
}
msg = AIMessage("", tool_calls=[tool_call])
config: RunnableConfig = {"configurable": {"__pregel_runtime": mock_runtime}}
result = node.invoke({"messages": [msg]}, config=config)
node.invoke({"messages": [msg]}, config=config)
assert result["messages"][-1].content == "ok"
assert captured["execution_info"] is exec_info
assert captured["execution_info"].thread_id == "t-1"
assert captured["execution_info"].task_id == "tk-1"
assert captured["server_info"] is server_info
assert captured["server_info"].assistant_id == "asst-1"
assert [tool.name for tool in captured["tools"]] == ["info_tool", "other_tool"]
async def test_tool_runtime_forwards_execution_info_server_info_and_tools_async() -> (
None
):
"""Test that execution_info, server_info, and tools are forwarded in async path."""
async def test_tool_runtime_forwards_execution_info_and_server_info_async() -> None:
"""Test that execution_info and server_info are forwarded in async path."""
from langgraph.runtime import ExecutionInfo, ServerInfo
exec_info = ExecutionInfo(
@@ -2100,15 +2090,9 @@ async def test_tool_runtime_forwards_execution_info_server_info_and_tools_async(
"""Async tool that captures runtime info."""
captured["execution_info"] = runtime.execution_info
captured["server_info"] = runtime.server_info
captured["tools"] = runtime.tools
return "ok"
@dec_tool
async def other_tool_async(y: int) -> str:
"""Another async tool available to the runtime."""
return str(y)
node = ToolNode([info_tool_async, other_tool_async])
node = ToolNode([info_tool_async])
tool_call = {
"name": "info_tool_async",
"args": {"x": 1},
@@ -2117,17 +2101,12 @@ async def test_tool_runtime_forwards_execution_info_server_info_and_tools_async(
}
msg = AIMessage("", tool_calls=[tool_call])
config: RunnableConfig = {"configurable": {"__pregel_runtime": mock_runtime}}
result = await node.ainvoke({"messages": [msg]}, config=config)
await node.ainvoke({"messages": [msg]}, config=config)
assert result["messages"][-1].content == "ok"
assert captured["execution_info"] is exec_info
assert captured["execution_info"].thread_id == "t-2"
assert captured["server_info"] is server_info
assert captured["server_info"].graph_id == "graph-2"
assert [tool.name for tool in captured["tools"]] == [
"info_tool_async",
"other_tool_async",
]
# --- InjectedToolArg security tests ---
+6 -6
View File
@@ -249,7 +249,7 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.0"
version = "1.3.0a2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
@@ -261,14 +261,14 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/92/fe/20190232d9b513242899dbb0c2bb77e31b4d61e343743adbe90ebc2603d2/langchain_core-1.3.0.tar.gz", hash = "sha256:14a39f528bf459aa3aa40d0a7f7f1bae7520d435ef991ae14a4ceb74d8c49046", size = 860755, upload-time = "2026-04-17T14:51:38.298Z" }
sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/f8/e2/dbfa347aa072a6dc4cd38d6f9ebfc730b4c14c258c47f480f4c5c546f177/langchain_core-1.3.0-py3-none-any.whl", hash = "sha256:baf16ee028475df177b9ab8869a751c79406d64a6f12125b93802991b566cced", size = 515140, upload-time = "2026-04-17T14:51:36.274Z" },
{ url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" },
]
[[package]]
name = "langgraph"
version = "1.1.9"
version = "1.1.7a2"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -281,7 +281,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = ">=1.3.0,<2" },
{ name = "langchain-core", specifier = "==1.3.0a2" },
{ name = "langgraph-checkpoint", editable = "../checkpoint" },
{ name = "langgraph-prebuilt", editable = "." },
{ name = "langgraph-sdk", editable = "../sdk-py" },
@@ -490,7 +490,7 @@ test = [
[[package]]
name = "langgraph-prebuilt"
version = "1.0.10"
version = "1.0.9"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
+6 -6
View File
@@ -262,7 +262,7 @@ wheels = [
[[package]]
name = "langchain-core"
version = "1.3.0"
version = "1.3.0a2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "jsonpatch" },
@@ -274,14 +274,14 @@ dependencies = [
{ name = "typing-extensions" },
{ name = "uuid-utils" },
]
sdist = { url = "https://files.pythonhosted.org/packages/92/fe/20190232d9b513242899dbb0c2bb77e31b4d61e343743adbe90ebc2603d2/langchain_core-1.3.0.tar.gz", hash = "sha256:14a39f528bf459aa3aa40d0a7f7f1bae7520d435ef991ae14a4ceb74d8c49046", size = 860755, upload-time = "2026-04-17T14:51:38.298Z" }
sdist = { url = "https://files.pythonhosted.org/packages/af/bc/0bff31fcaff174d86031cc713471a3e85ed4ec8e5cd95ad0217f2aced20e/langchain_core-1.3.0a2.tar.gz", hash = "sha256:52d978c84552b74b9a3f16c1fced84f9e27cc96d7a67c601925ce6cbc4ea3cf9", size = 854580, upload-time = "2026-04-13T14:37:55.745Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/f8/e2/dbfa347aa072a6dc4cd38d6f9ebfc730b4c14c258c47f480f4c5c546f177/langchain_core-1.3.0-py3-none-any.whl", hash = "sha256:baf16ee028475df177b9ab8869a751c79406d64a6f12125b93802991b566cced", size = 515140, upload-time = "2026-04-17T14:51:36.274Z" },
{ url = "https://files.pythonhosted.org/packages/0e/14/03c09686602567059f26af29de0c44546a83af2f2aa29925e61040e43ea2/langchain_core-1.3.0a2-py3-none-any.whl", hash = "sha256:9e929a34f0b0c6c1255e395a1de34f8626893ceb4cdae550a22a0bd18c87be54", size = 510233, upload-time = "2026-04-13T14:37:54.277Z" },
]
[[package]]
name = "langgraph"
version = "1.1.9"
version = "1.1.7a2"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -294,7 +294,7 @@ dependencies = [
[package.metadata]
requires-dist = [
{ name = "langchain-core", specifier = ">=1.3.0,<2" },
{ name = "langchain-core", specifier = "==1.3.0a2" },
{ name = "langgraph-checkpoint", editable = "../checkpoint" },
{ name = "langgraph-prebuilt", editable = "../prebuilt" },
{ name = "langgraph-sdk", editable = "." },
@@ -413,7 +413,7 @@ test = [
[[package]]
name = "langgraph-prebuilt"
version = "1.0.10"
version = "1.0.9"
source = { editable = "../prebuilt" }
dependencies = [
{ name = "langchain-core" },
-133
View File
@@ -1,133 +0,0 @@
# feat(channels): DeltaChannel — O(N) incremental checkpoint storage
## The problem
LangGraph checkpoints store the **full accumulated value** of every channel on every step. For a `messages` channel backed by `add_messages`, that means each checkpoint blob contains the entire conversation history up to that point.
Storage cost grows **O(N²)** in the number of turns:
| Step | Checkpoint blob |
|------|----------------|
| 1 | [msg_1] |
| 2 | [msg_1, msg_2] |
| N | [msg_1, ..., msg_N] |
At 100K tokens of conversation data, a single thread accumulates ~250 MB; with large messages or file attachments costs scale even faster.
## The fix: `DeltaChannel`
`DeltaChannel` is an opt-in wrapper around any binary reducer that stores only a **sentinel marker** in `checkpoint_blobs` rather than the full accumulated value. The actual per-step writes stay in `checkpoint_writes` (which every checkpointer already writes unconditionally). At read time the saver walks the ancestor chain, collects all writes for the channel, and replays them through the reducer.
Storage scales **O(N)** — the sentinel blob is effectively zero bytes, and the writes table already exists.
```python
from langgraph.channels.delta import DeltaChannel
from langgraph.graph.message import add_messages
class State(TypedDict):
# Before: O(N²) storage
messages: Annotated[list[AnyMessage], add_messages]
# After: O(N) storage
messages: Annotated[list[AnyMessage], DeltaChannel(add_messages)]
```
## Benchmarks
Simulated with realistic paragraph-length messages (~100 tokens each, ~400 chars). Each turn = one human + one AI message (~200 tokens total).
### Storage (InMemorySaver)
| turns | ctx | add_msgs | delta | savings |
|------:|----:|---------:|------:|--------:|
| 10 | ~2K tok | 108.6 KB | 4.0 KB | 27x |
| 25 | ~5K tok | 649.0 KB | 10.1 KB | 64x |
| 50 | ~10K tok | 2.6 MB | 20.2 KB | 126x |
| 100 | ~20K tok | 10.2 MB | 40.5 KB | 251x |
| 500 | ~100K tok | 252.6 MB | 202.8 KB | 1245x |
Savings grow with N because `add_messages` is O(N²) while `DeltaChannel` is O(N). The sentinel blob itself is essentially zero bytes.
### Read latency (avg of 5 `get_state` calls = cost per `invoke`)
| turns | ctx | add_msgs | delta |
|------:|----:|---------:|------:|
| 10 | ~2K tok | 0.1ms | 0.2ms |
| 25 | ~5K tok | 0.2ms | 0.5ms |
| 50 | ~10K tok | 0.4ms | 1.5ms |
| 100 | ~20K tok | 0.7ms | 4.8ms |
| 500 | ~100K tok | 5.8ms | 114.9ms |
**This cost is paid once per `invoke`/`stream` call, not per node.** Within a single invocation, all channels are loaded into memory once at the start and shared across every node — there is no per-node reconstruction. The 114.9ms at 500 turns is what you pay each time a user sends a new message, not on each step of the graph.
## How it works
**Write:** `DeltaChannel.checkpoint()` always emits `DeltaChannelSentinel()` — a tiny marker (zero payload bytes) stored in `checkpoint_blobs`. Per-step writes flow into `checkpoint_writes` as they normally do for every channel.
**Read:** The saver detects `DeltaChannelSentinel` values in `channel_values` and replaces them by calling `get_channel_writes` / `aget_channel_writes`, which walks the ancestor checkpoint chain and collects all writes for that channel (oldest→newest). `DeltaChannel.from_checkpoint()` replays those writes through the operator to reconstruct the full value.
**Saver implementations:**
- `InMemorySaver` — direct dict traversal of `self.storage` and `self.writes`, no I/O
- `PostgresSaver` (sync + async) — two queries: one cheap ID walk across the thread, one `ANY()` fetch of writes; no recursive CTE
- All other savers — `BaseCheckpointSaver.get_channel_writes` fallback via `list()`, with a re-entrancy guard to prevent infinite recursion
## Changes
**`libs/checkpoint`**
- `base/__init__.py` — add `DeltaChannelSentinel` marker dataclass; add `get_channel_writes` / `aget_channel_writes` to `BaseCheckpointSaver` with a `list()`-based fallback and re-entrancy guard
**`libs/checkpoint/memory`**
- `memory/__init__.py` — `get_channel_writes` via direct dict traversal; `_resolve_delta_channels` helper called in `get_tuple` / `aget_tuple` to replace sentinels with reconstructed write lists
**`libs/langgraph`**
- `channels/delta.py` — `DeltaChannel` implementation: `checkpoint()` always emits sentinel, `from_checkpoint()` replays writes list
- `channels/__init__.py` — export `DeltaChannel`
- `graph/state.py` — recognize `DeltaChannel` as a valid channel annotation
- `pregel/_checkpoint.py` / `pregel/_loop.py` — wire `after_checkpoint` hook; call it after each checkpointing step so `DeltaChannel` can advance internal state
**`libs/checkpoint-postgres`**
- `postgres/base.py` — `_get_channel_writes_cur` two-query ancestor walk (sync); `_resolve_delta_channels` called after `_load_blobs`
- `postgres/aio.py` — `_aget_channel_writes_cur` (async counterpart)
## Open questions
**Should we add a compile-time capability check?**
Currently misconfiguring `DeltaChannel` with an unsupported saver only errors at runtime on first reload. A protocol-based check at `compile()` time would give an early warning without requiring a manual boolean flag.
**`snapshot_every` for bounded reconstruction cost?**
Both per-invoke read latency and total write wall time grow O(N) per invoke / O(N²) total as the conversation lengthens. A `snapshot_every` parameter — periodically store a full snapshot in `checkpoint_blobs` to cap chain depth — would bound reconstruction cost and is a natural follow-up once the core design is stable.
## Backwards compatibility
| Scenario | Behaviour |
|----------|-----------|
| Existing graph using `add_messages` | Unaffected — no code or schema changes |
| `DeltaChannel` loading an old full-list checkpoint blob | Handled via backwards-compat path in `from_checkpoint` |
| `DeltaChannel` with `InMemorySaver` or `PostgresSaver` | Fully supported |
| Time-travel to a past checkpoint | Ancestor walk uses the version at that checkpoint — correct by construction |
| `Overwrite` value | Resets the effective chain; reconstruction starts from that step |
## Test plan
- [x] `DeltaChannel` unit tests: `update` → `checkpoint` lifecycle, `from_checkpoint` chain replay, backwards-compat with plain list, `Overwrite` resets chain
- [x] `InMemorySaver` `get_channel_writes`: assembles write list from dict storage
- [x] Serde round-trip for `DeltaChannelSentinel`
- [x] End-to-end graph tests: multi-turn conversations accumulate correctly, time-travel reconstructs correct partial history
- [x] `PostgresSaver` two-query chain reconstruction (sync + async)
- [x] `BaseCheckpointSaver` fallback path via `list()` with re-entrancy guard
- [x] Storage benchmark: `DeltaChannel` uses strictly less storage than `add_messages` at all measured turn counts
---
## Changes from previous base branch
The previous version stored `DeltaValue` objects (containing the per-step writes) directly in `checkpoint_blobs` and used a `DeltaChainValue` to represent the assembled chain. Reconstruction required a dedicated `get_delta_chain` / `aget_delta_chain` protocol and a recursive CTE in Postgres.
This version pivots to a simpler design:
- **Sentinel in blobs, writes in `checkpoint_writes`** — `checkpoint_blobs` stores only a zero-byte `DeltaChannelSentinel` marker. The actual per-step data already lives in `checkpoint_writes` (written unconditionally by every checkpointer), so blob storage is essentially free. This is why storage savings jump to 1245x at 500 turns.
- **No custom serde type for the delta payload** — `DeltaValue` / `DeltaChainValue` and the `"delta"` serde type tag are gone. Writes are deserialized with the same serde path they were originally written with.
- **Postgres: two queries instead of a recursive CTE** — fetch all `(checkpoint_id, parent_checkpoint_id)` pairs for the thread, walk the ancestor chain in Python, then fetch writes with a plain `ANY()` filter.
- **Universal fallback on `BaseCheckpointSaver`** — the base class now provides `get_channel_writes` via `list()`, so any third-party saver works without modification.
- **`snapshot_every` removed** — deferred as a follow-up; the simpler design is easier to reason about and delivers larger storage savings.