Compare commits

..
Author SHA1 Message Date
Will Fu-Hinthorn 685a755baf chore: use bounded queue for checkpointing 2026-04-09 17:33:50 -07:00
Will Fu-HinthornandClaude Opus 4.6 bde0c47cd9 fix: handle non-deterministic ordering in test_imp_nested
Nested tasks run concurrently so output order is non-deterministic.
Sort intermediate results before comparison, matching the pattern
already used by test_imp_task.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-08 18:04:55 -07:00
Will Fu-Hinthorn 7b034c2d25 chore: update conformance lint 2026-04-08 16:02:42 -07:00
6242b99e06 chore(checkpoint-conformance): remove test_list_global_search, bump to 0.0.2 (#7444)
## Summary
- Remove `test_list_global_search` from the conformance test suite. This
test required cross-thread `alist(None, filter=...)` support that not
all checkpointer implementations provide.
- Remove the corresponding entry from `ALL_LIST_TESTS`.
- Bump `langgraph-checkpoint-conformance` version from 0.0.1 to 0.0.2.

## Test plan
- [x] Verify `test_list_global_search` function definition is fully
removed
- [x] Verify `test_list_global_search` is removed from `ALL_LIST_TESTS`
- [x] Verify version bumped to 0.0.2 in pyproject.toml

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-authored-by: Will Fu-Hinthorn <will@langchain.dev>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-08 05:54:37 -07:00
890147681d chore(cli): add validate command (#7438)
add validate command

---------

Co-authored-by: Will Fu-Hinthorn <will@langchain.dev>
2026-04-07 18:29:24 -07:00
12 changed files with 613 additions and 162 deletions
@@ -339,37 +339,6 @@ async def test_list_metadata_custom_keys(
assert results[0].metadata["run_id"] == "run-abc"
async def test_list_global_search(
saver: BaseCheckpointSaver,
) -> None:
"""alist(None, filter=...) searches across all threads."""
tid1, tid2 = str(uuid4()), str(uuid4())
# Use a unique marker so we don't collide with other tests' data
marker = str(uuid4())
cfg1 = generate_config(tid1)
cp1 = generate_checkpoint()
await saver.aput(cfg1, cp1, generate_metadata(source="input", marker=marker), {})
cfg2 = generate_config(tid2)
cp2 = generate_checkpoint()
await saver.aput(cfg2, cp2, generate_metadata(source="loop", marker=marker), {})
# Search across all threads with filter
results = []
async for tup in saver.alist(None, filter={"source": "input", "marker": marker}):
results.append(tup)
assert len(results) == 1
assert results[0].config["configurable"]["thread_id"] == tid1
# Search with marker only — should find both
results = []
async for tup in saver.alist(None, filter={"marker": marker}):
results.append(tup)
assert len(results) == 2
ALL_LIST_TESTS = [
test_list_all,
test_list_by_thread,
@@ -380,7 +349,6 @@ ALL_LIST_TESTS = [
test_list_metadata_filter_multiple_keys,
test_list_metadata_filter_no_match,
test_list_metadata_custom_keys,
test_list_global_search,
test_list_before,
test_list_limit,
test_list_limit_plus_before,
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint-conformance"
version = "0.0.1"
version = "0.0.2"
description = "Conformance test suite for LangGraph checkpointer implementations."
authors = [{name = "William FH", email = "13333726+hinthornw@users.noreply.github.com"}]
requires-python = ">=3.10"
+1 -1
View File
@@ -263,7 +263,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint-conformance"
version = "0.0.1"
version = "0.0.2"
source = { editable = "." }
dependencies = [
{ name = "langgraph-checkpoint" },
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.20"
__version__ = "0.4.21"
+42
View File
@@ -817,6 +817,48 @@ def dev(
)
# ---------------------------------------------------------------------------
# validate command
# ---------------------------------------------------------------------------
@OPT_CONFIG
@cli.command(help="✅ Validate the LangGraph configuration file.")
@log_command
def validate(config: pathlib.Path):
import json
try:
with open(config) as f:
raw_config = json.load(f)
except json.JSONDecodeError as e:
raise click.UsageError(f"Invalid JSON in {config}: {e.args[0]}") from None
# Check for unknown keys before validation so they show alongside any error.
unknown_warnings = langgraph_cli.config.get_unknown_keys(raw_config)
try:
config_json = langgraph_cli.config.validate_config_file(config)
except (click.UsageError, ValueError) as e:
click.secho(f"Error: {e}", fg="red", err=True)
if unknown_warnings:
click.echo(err=True)
for warning in unknown_warnings:
click.secho(f" warning: {warning}", fg="yellow", err=True)
raise SystemExit(1) from None
num_graphs = len(config_json.get("graphs", {}))
click.secho(
f"Configuration file {config} is valid. "
f"({num_graphs} graph{'s' if num_graphs != 1 else ''} found)",
fg="green",
)
if unknown_warnings:
click.echo()
for warning in unknown_warnings:
click.secho(f" warning: {warning}", fg="yellow")
# ---------------------------------------------------------------------------
# new command
# ---------------------------------------------------------------------------
+77 -17
View File
@@ -182,7 +182,11 @@ def validate_config(config: Config) -> Config:
"Version must be major or major.minor or major.minor.patch."
)
except TypeError:
raise click.UsageError(f"Invalid version format: {api_version}") from None
raise click.UsageError(
f"Invalid version format: {api_version}.\n\n"
"Pin to a minor version, e.g.:\n"
' "api_version": "0.8"'
) from None
config = {
"node_version": node_version,
@@ -220,45 +224,51 @@ def validate_config(config: Config) -> Config:
if major < min_major:
raise click.UsageError(
f"Node.js version {node_version} is not supported. "
f"Minimum required version is {MIN_NODE_VERSION}."
f"Minimum required version is {MIN_NODE_VERSION}.\n\n"
f"Set node_version to {MIN_NODE_VERSION} or higher:\n"
f' "node_version": "{MIN_NODE_VERSION}"'
)
except ValueError as e:
raise click.UsageError(str(e)) from None
if pip_installer := config.get("pip_installer"):
if pip_installer == "uv_lock":
raise click.UsageError(
"pip_installer 'uv_lock' has been replaced. Use "
'`source: {"kind": "uv", "root": "..", '
'"package": "my-agent"}`.'
)
if pip_installer not in ["auto", "pip", "uv"]:
raise click.UsageError(
f"Invalid pip_installer: '{pip_installer}'. "
"Must be 'auto', 'pip', or 'uv'."
"Consider using uv-based source management instead:\n\n"
' "source": {"kind": "uv", "root": ".."}'
)
source = config.get("source")
source_kind = _get_source_kind(config)
if source is not None and not isinstance(source, dict):
raise click.UsageError("`source` must be an object.")
raise click.UsageError(
"`source` must be an object, e.g.:\n"
' "source": {"kind": "uv", "root": ".."}'
)
if source is not None and source_kind != "uv":
raise click.UsageError("Invalid source.kind. Supported values: 'uv'.")
raise click.UsageError(
"Invalid source.kind. The only supported value is 'uv':\n"
' "source": {"kind": "uv", "root": ".."}'
)
if config.get("python_version"):
pyversion = config["python_version"]
if not pyversion.count(".") == 1 or not all(
part.isdigit() for part in pyversion.split("-")[0].split(".")
):
parts = pyversion.split("-")[0].split(".")
fix = f"{parts[0]}.{parts[1]}" if len(parts) >= 2 else MIN_PYTHON_VERSION
raise click.UsageError(
f"Invalid Python version format: {pyversion}. "
"Use 'major.minor' format (e.g., '3.11'). "
"Patch version cannot be specified."
"Use 'major.minor' format — patch version cannot be specified.\n\n"
f' "python_version": "{fix}"'
)
if _parse_version(pyversion) < _parse_version(MIN_PYTHON_VERSION):
raise click.UsageError(
f"Python version {pyversion} is not supported. "
f"Minimum required version is {MIN_PYTHON_VERSION}."
f"Minimum required version is {MIN_PYTHON_VERSION}.\n\n"
f' "python_version": "{MIN_PYTHON_VERSION}"'
)
if "bullseye" in pyversion:
raise click.UsageError(
@@ -269,12 +279,16 @@ def validate_config(config: Config) -> Config:
if source_kind != "uv" and not config["dependencies"]:
raise click.UsageError(
"No dependencies found in config. "
"Add at least one dependency to 'dependencies' list."
"Consider using uv-based source management:\n\n"
' "source": {"kind": "uv", "root": ".."}'
)
if not config.get("graphs"):
raise click.UsageError(
"No graphs found in config. Add at least one graph to 'graphs' dictionary."
"No graphs found in config. Add at least one graph, e.g.:\n"
' "graphs": {\n'
' "agent": "./my_agent/graph.py:graph"\n'
" }"
)
# Validate image_distro config
@@ -287,7 +301,8 @@ def validate_config(config: Config) -> Config:
if image_distro not in Distros.__args__:
raise click.UsageError(
f"Invalid image_distro: '{image_distro}'. "
"Must be one of 'debian', 'wolfi', or 'bookworm'."
f"Must be one of: {', '.join(repr(d) for d in Distros.__args__)}.\n\n"
' "image_distro": "wolfi" (recommended)'
)
if source_kind == "uv":
@@ -369,6 +384,51 @@ def validate_config(config: Config) -> Config:
return config
# Keys recognized by validate_config (used to detect unknown fields).
_KNOWN_CONFIG_KEYS = {
"python_version",
"node_version",
"api_version",
"base_image",
"image_distro",
"pip_config_file",
"pip_installer",
"source",
"dependencies",
"dockerfile_lines",
"graphs",
"env",
"store",
"auth",
"encryption",
"http",
"webhooks",
"checkpointer",
"ui",
"ui_config",
"keep_pkg_tools",
# Internal / legacy (still recognized, may error separately)
"_INTERNAL_docker_tag",
"project_root",
"package",
}
def get_unknown_keys(raw_config: dict) -> list[str]:
"""Return warnings for unrecognized top-level keys (typos, etc.)."""
import difflib
unknown = set(raw_config) - _KNOWN_CONFIG_KEYS
warnings: list[str] = []
for key in sorted(unknown):
close = difflib.get_close_matches(key, _KNOWN_CONFIG_KEYS, n=1)
if close:
warnings.append(f"Unknown key '{key}' — did you mean '{close[0]}'?")
else:
warnings.append(f"Unknown key '{key}' is not a recognized config field.")
return warnings
def validate_config_file(config_path: pathlib.Path) -> Config:
"""Load and validate a configuration file."""
with open(config_path) as f:
+2 -2
View File
@@ -404,7 +404,7 @@ def test_validate_config_pip_installer():
}
)
assert "Invalid pip_installer: 'conda'" in str(exc_info.value)
assert "Must be 'auto', 'pip', or 'uv'" in str(exc_info.value)
assert "uv-based source management" in str(exc_info.value)
with pytest.raises(click.UsageError) as exc_info:
validate_config(
@@ -417,7 +417,7 @@ def test_validate_config_pip_installer():
)
assert "Invalid pip_installer: 'invalid'" in str(exc_info.value)
with pytest.raises(click.UsageError, match="has been replaced"):
with pytest.raises(click.UsageError, match="Invalid pip_installer: 'uv_lock'"):
validate_config(
{
"python_version": "3.11",
@@ -0,0 +1,215 @@
from __future__ import annotations
import asyncio
import os
import queue
import threading
from collections.abc import Callable
from contextlib import AbstractAsyncContextManager, AbstractContextManager
from dataclasses import dataclass
from types import TracebackType
from typing import Any
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
ChannelVersions,
Checkpoint,
CheckpointMetadata,
)
QUEUE_PUT_TIMEOUT = 0.05
CHECKPOINT_BACKLOG_ENV_VAR = "LANGGRAPH_CHECKPOINT_BACKLOG"
DEFAULT_CHECKPOINT_BACKLOG = 10
@dataclass(frozen=True)
class CheckpointRequest:
config: RunnableConfig
checkpoint: Checkpoint
metadata: CheckpointMetadata
new_versions: ChannelVersions
def _raise(error: BaseException) -> None:
raise error
def resolve_checkpoint_backlog() -> int:
if raw := os.getenv(CHECKPOINT_BACKLOG_ENV_VAR):
try:
backlog = int(raw)
except ValueError:
return DEFAULT_CHECKPOINT_BACKLOG
if backlog > 0:
return backlog
return DEFAULT_CHECKPOINT_BACKLOG
class SyncCheckpointWriter(AbstractContextManager):
def __init__(
self,
put: Callable[
[RunnableConfig, Checkpoint, CheckpointMetadata, ChannelVersions], Any
],
*,
max_pending: int | None = None,
) -> None:
self.put = put
max_pending = (
resolve_checkpoint_backlog() if max_pending is None else max_pending
)
self.queue: queue.Queue[CheckpointRequest | None] = queue.Queue(max_pending)
self.error: BaseException | None = None
self.closed = False
self.thread = threading.Thread(
target=self._run,
name="langgraph-checkpoint-writer",
daemon=True,
)
def __enter__(self) -> SyncCheckpointWriter:
self.thread.start()
return self
def submit(self, request: CheckpointRequest) -> None:
self._ensure_open()
while True:
self._raise_if_broken()
try:
self.queue.put(request, timeout=QUEUE_PUT_TIMEOUT)
except queue.Full:
continue
else:
self._raise_if_broken()
return
def _run(self) -> None:
while True:
item = self.queue.get()
if item is None:
return
try:
self.put(
item.config,
item.checkpoint,
item.metadata,
item.new_versions,
)
except BaseException as exc:
self.error = exc
return
def _ensure_open(self) -> None:
if self.closed:
raise RuntimeError("Checkpoint writer is closed")
def _raise_if_broken(self) -> None:
if self.error is not None:
_raise(self.error)
def __exit__(
self,
exc_type: type[BaseException] | None,
exc_value: BaseException | None,
traceback: TracebackType | None,
) -> bool | None:
self.closed = True
while self.thread.is_alive():
if self.error is not None:
break
try:
self.queue.put(None, timeout=QUEUE_PUT_TIMEOUT)
except queue.Full:
continue
else:
break
self.thread.join()
if exc_type is None and self.error is not None:
_raise(self.error)
return None
class AsyncCheckpointWriter(AbstractAsyncContextManager):
def __init__(
self,
put: Callable[
[RunnableConfig, Checkpoint, CheckpointMetadata, ChannelVersions], Any
],
*,
max_pending: int | None = None,
) -> None:
self.put = put
max_pending = (
resolve_checkpoint_backlog() if max_pending is None else max_pending
)
self.queue: asyncio.Queue[CheckpointRequest | None] = asyncio.Queue(max_pending)
self.error: BaseException | None = None
self.closed = False
self.task: asyncio.Task[None] | None = None
async def __aenter__(self) -> AsyncCheckpointWriter:
self.task = asyncio.create_task(self._run(), name="langgraph-checkpoint-writer")
return self
async def submit(self, request: CheckpointRequest) -> None:
self._ensure_open()
while True:
self._raise_if_broken()
try:
await asyncio.wait_for(
self.queue.put(request),
timeout=QUEUE_PUT_TIMEOUT,
)
except asyncio.TimeoutError:
continue
else:
self._raise_if_broken()
return
async def _run(self) -> None:
while True:
item = await self.queue.get()
if item is None:
return
try:
await self.put(
item.config,
item.checkpoint,
item.metadata,
item.new_versions,
)
except BaseException as exc:
self.error = exc
return
def _ensure_open(self) -> None:
if self.closed:
raise RuntimeError("Checkpoint writer is closed")
def _raise_if_broken(self) -> None:
if self.error is not None:
_raise(self.error)
async def __aexit__(
self,
exc_type: type[BaseException] | None,
exc_value: BaseException | None,
traceback: TracebackType | None,
) -> None:
self.closed = True
while self.task is not None and not self.task.done():
if self.error is not None:
break
try:
await asyncio.wait_for(
self.queue.put(None),
timeout=QUEUE_PUT_TIMEOUT,
)
except asyncio.TimeoutError:
continue
else:
break
if self.task is not None:
await self.task
if exc_type is None and self.error is not None:
_raise(self.error)
+109 -90
View File
@@ -2,7 +2,6 @@ from __future__ import annotations
import asyncio
import binascii
import concurrent.futures
from collections import defaultdict, deque
from collections.abc import Callable, Iterator, Mapping, Sequence
from contextlib import (
@@ -27,7 +26,6 @@ from langgraph.cache.base import BaseCache
from langgraph.checkpoint.base import (
WRITES_IDX_MAP,
BaseCheckpointSaver,
ChannelVersions,
Checkpoint,
CheckpointMetadata,
CheckpointTuple,
@@ -92,6 +90,11 @@ from langgraph.pregel._checkpoint import (
create_checkpoint,
empty_checkpoint,
)
from langgraph.pregel._checkpoint_writer import (
AsyncCheckpointWriter,
CheckpointRequest,
SyncCheckpointWriter,
)
from langgraph.pregel._executor import (
AsyncBackgroundExecutor,
BackgroundExecutor,
@@ -166,19 +169,6 @@ class PregelLoop:
checkpointer_get_next_version: GetNextVersion
checkpointer_put_writes: Callable[[RunnableConfig, WritesT, str], Any] | None
checkpointer_put_writes_accepts_task_path: bool
_checkpointer_put_after_previous: (
Callable[
[
concurrent.futures.Future | None,
RunnableConfig,
Checkpoint,
str,
ChannelVersions,
],
Any,
]
| None
)
_migrate_checkpoint: Callable[[Checkpoint], None] | None
submit: Submit
channels: Mapping[str, BaseChannel]
@@ -491,7 +481,7 @@ class PregelLoop:
)
# produce debug output
if self._checkpointer_put_after_previous is not None:
if self.checkpointer is not None:
self._emit(
"checkpoints",
map_debug_checkpoint,
@@ -537,7 +527,7 @@ class PregelLoop:
return True
def after_tick(self) -> None:
def _after_tick(self) -> CheckpointRequest | None:
# finish superstep
writes = [w for t in self.tasks.values() for w in t.writes]
# all tasks have finished
@@ -562,14 +552,14 @@ class PregelLoop:
# only replay (re-execute) done tasks on the first tick
self.is_replaying = False
# save checkpoint
self._put_checkpoint({"source": "loop"})
# after execution, check if we should interrupt
return self._prepare_checkpoint({"source": "loop"})
def _finish_after_tick(self) -> None:
if self.interrupt_after and should_interrupt(
self.checkpoint, self.interrupt_after, self.tasks.values()
):
self.status = "interrupt_after"
raise GraphInterrupt()
# unset resuming flag
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
@@ -619,7 +609,7 @@ class PregelLoop:
def _first(
self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None
) -> set[str] | None:
) -> tuple[set[str] | None, CheckpointRequest | None]:
# Resuming from a previous checkpoint requires two things:
# 1. A prior checkpoint exists (channel_versions is non-empty)
# 2. The input signals continuation (not a fresh run with new input)
@@ -713,6 +703,7 @@ class PregelLoop:
)
if updated_channels is not None:
updated_channels.update(null_updated_channels)
checkpoint_request = None
# proceed past previous checkpoint
if is_resuming:
self.checkpoint["versions_seen"].setdefault(INTERRUPT, {})
@@ -755,7 +746,7 @@ class PregelLoop:
)
# save input checkpoint
self.updated_channels = updated_channels
self._put_checkpoint({"source": "input"})
checkpoint_request = self._prepare_checkpoint({"source": "input"})
elif CONFIG_KEY_RESUMING not in configurable:
raise EmptyInputError(f"Received no input for {input_keys}")
# Propagate resuming and replaying flags to subgraphs.
@@ -785,9 +776,11 @@ class PregelLoop:
)
# set flag
self.status = "pending"
return updated_channels
return updated_channels, checkpoint_request
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
def _prepare_checkpoint(
self, metadata: CheckpointMetadata
) -> CheckpointRequest | None:
# assign step and parents
exiting = metadata is self.checkpoint_metadata
if exiting and self.checkpoint["id"] == self.checkpoint_id_saved:
@@ -797,8 +790,7 @@ class PregelLoop:
metadata["step"] = self.step
metadata["parents"] = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
self.checkpoint_metadata = metadata
# do checkpoint?
do_checkpoint = self._checkpointer_put_after_previous is not None and (
do_checkpoint = self.checkpointer is not None and (
exiting or self.durability != "exit"
)
# create new checkpoint
@@ -820,9 +812,8 @@ class PregelLoop:
for value in self.checkpoint["channel_values"][TASKS]
]
self.checkpoint["channel_values"][TASKS] = sanitized_tasks
# bail if no checkpointer
if do_checkpoint and self._checkpointer_put_after_previous is not None:
request = None
if do_checkpoint:
self.prev_checkpoint_config = (
self.checkpoint_config
if CONFIG_KEY_CHECKPOINT_ID in self.checkpoint_config[CONF]
@@ -844,17 +835,11 @@ class PregelLoop:
self.checkpoint_previous_versions, channel_versions
)
self.checkpoint_previous_versions = channel_versions
# save it, without blocking
# if there's a previous checkpoint save in progress, wait for it
# ensuring checkpointers receive checkpoints in order
self._put_checkpoint_fut = self.submit(
self._checkpointer_put_after_previous,
getattr(self, "_put_checkpoint_fut", None),
self.checkpoint_config,
copy_checkpoint(self.checkpoint),
self.checkpoint_metadata,
new_versions,
request = CheckpointRequest(
config=self.checkpoint_config,
checkpoint=copy_checkpoint(self.checkpoint),
metadata=self.checkpoint_metadata,
new_versions=new_versions,
)
self.checkpoint_config = {
**self.checkpoint_config,
@@ -866,28 +851,18 @@ class PregelLoop:
if not exiting:
# increment step
self.step += 1
return request
def _suppress_interrupt(
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
raise NotImplementedError
def _finalize_suppress(
self,
exc_type: type[BaseException] | None,
exc_value: BaseException | None,
traceback: TracebackType | None,
) -> bool | None:
# persist current checkpoint and writes
if self.durability == "exit" and (
# if it's a top graph
not self.is_nested
# or a nested graph with error or interrupt
or exc_value is not None
# or a nested graph with checkpointer=True
or all(NS_END not in part for part in self.checkpoint_ns)
):
self._put_checkpoint(self.checkpoint_metadata)
self._put_pending_writes()
# suppress interrupt
suppress = isinstance(exc_value, GraphInterrupt) and not self.is_nested
if suppress:
# emit one last "values" event, with pending writes applied
if (
hasattr(self, "tasks")
and self.checkpoint_pending_writes
@@ -912,7 +887,6 @@ class PregelLoop:
[w for t in self.tasks.values() for w in t.writes],
self.channels,
)
# emit INTERRUPT if exception is empty (otherwise emitted by put_writes)
if exc_value is not None and (not exc_value.args or not exc_value.args[0]):
self._emit(
"updates",
@@ -920,13 +894,26 @@ class PregelLoop:
[{INTERRUPT: cast(GraphInterrupt, exc_value).args[0]}]
),
)
# save final output
self.output = read_channels(self.channels, self.output_keys)
# suppress interrupt
return True
elif exc_type is None:
# save final output
self.output = read_channels(self.channels, self.output_keys)
return None
def _suppress_interrupt(
self,
exc_type: type[BaseException] | None,
exc_value: BaseException | None,
traceback: TracebackType | None,
) -> bool | None:
if self.durability == "exit" and (
not self.is_nested
or exc_value is not None
or all(NS_END not in part for part in self.checkpoint_ns)
):
self._put_checkpoint(self.checkpoint_metadata)
self._put_pending_writes()
return self._finalize_suppress(exc_type, exc_value)
def _emit(
self,
@@ -1072,26 +1059,30 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
)
else:
self.checkpointer_get_next_version = increment
self._checkpointer_put_after_previous = None # type: ignore[assignment]
self.checkpointer_put_writes = None
self.checkpointer_put_writes_accepts_task_path = False
self._checkpoint_writer: SyncCheckpointWriter | None = None
def _checkpointer_put_after_previous(
self,
prev: concurrent.futures.Future | None,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
try:
if prev is not None:
prev.result()
finally:
def _dispatch_checkpoint_request(self, request: CheckpointRequest) -> None:
if self.durability == "async" and self._checkpoint_writer is not None:
self._checkpoint_writer.submit(request)
else:
cast(BaseCheckpointSaver, self.checkpointer).put(
config, checkpoint, metadata, new_versions
request.config,
request.checkpoint,
request.metadata,
request.new_versions,
)
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
if request := self._prepare_checkpoint(metadata):
self._dispatch_checkpoint_request(request)
def after_tick(self) -> None:
if request := self._after_tick():
self._dispatch_checkpoint_request(request)
self._finish_after_tick()
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
return ()
@@ -1186,6 +1177,10 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
else []
)
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
if self.checkpointer is not None and self.durability == "async":
self._checkpoint_writer = self.stack.enter_context(
SyncCheckpointWriter(self.checkpointer.put)
)
self.channels, self.managed = channels_from_checkpoint(
self.specs, self.checkpoint
)
@@ -1194,12 +1189,14 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
self.step = self.checkpoint_metadata["step"] + 1
self.stop = self.step + self.config["recursion_limit"] + 1
self.checkpoint_previous_versions = self.checkpoint["channel_versions"].copy()
self.updated_channels = self._first(
self.updated_channels, checkpoint_request = self._first(
input_keys=self.input_keys,
updated_channels=set(self.checkpoint.get("updated_channels")) # type: ignore[arg-type]
if self.checkpoint.get("updated_channels")
else None,
)
if checkpoint_request is not None:
self._dispatch_checkpoint_request(checkpoint_request)
return self
@@ -1268,26 +1265,26 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
)
else:
self.checkpointer_get_next_version = increment
self._checkpointer_put_after_previous = None # type: ignore[assignment]
self.checkpointer_put_writes = None
self.checkpointer_put_writes_accepts_task_path = False
self._checkpoint_writer: AsyncCheckpointWriter | None = None
async def _checkpointer_put_after_previous(
self,
prev: asyncio.Task | None,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
try:
if prev is not None:
await prev
finally:
async def _dispatch_checkpoint_request(self, request: CheckpointRequest) -> None:
if self.durability == "async" and self._checkpoint_writer is not None:
await self._checkpoint_writer.submit(request)
else:
await cast(BaseCheckpointSaver, self.checkpointer).aput(
config, checkpoint, metadata, new_versions
request.config,
request.checkpoint,
request.metadata,
request.new_versions,
)
async def aafter_tick(self) -> None:
if request := self._after_tick():
await self._dispatch_checkpoint_request(request)
self._finish_after_tick()
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
return []
@@ -1332,6 +1329,22 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
},
)
async def _asuppress_interrupt(
self,
exc_type: type[BaseException] | None,
exc_value: BaseException | None,
traceback: TracebackType | None,
) -> bool | None:
if self.durability == "exit" and (
not self.is_nested
or exc_value is not None
or all(NS_END not in part for part in self.checkpoint_ns)
):
if request := self._prepare_checkpoint(self.checkpoint_metadata):
await self._dispatch_checkpoint_request(request)
self._put_pending_writes()
return self._finalize_suppress(exc_type, exc_value)
# context manager
async def __aenter__(self) -> Self:
@@ -1387,20 +1400,26 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
self.submit = await self.stack.enter_async_context(
AsyncBackgroundExecutor(self.config)
)
if self.checkpointer is not None and self.durability == "async":
self._checkpoint_writer = await self.stack.enter_async_context(
AsyncCheckpointWriter(self.checkpointer.aput)
)
self.channels, self.managed = channels_from_checkpoint(
self.specs, self.checkpoint
)
self.stack.push(self._suppress_interrupt)
self.stack.push_async_exit(self._asuppress_interrupt)
self.status = "input"
self.step = self.checkpoint_metadata["step"] + 1
self.stop = self.step + self.config["recursion_limit"] + 1
self.checkpoint_previous_versions = self.checkpoint["channel_versions"].copy()
self.updated_channels = self._first(
self.updated_channels, checkpoint_request = self._first(
input_keys=self.input_keys,
updated_channels=set(self.checkpoint.get("updated_channels")) # type: ignore[arg-type]
if self.checkpoint.get("updated_channels")
else None,
)
if checkpoint_request is not None:
await self._dispatch_checkpoint_request(checkpoint_request)
return self
+1 -7
View File
@@ -2751,9 +2751,6 @@ class Pregel(
_state_mapper,
)
loop.after_tick()
# wait for checkpoint
if durability_ == "sync":
loop._put_checkpoint_fut.result()
# emit output
yield from _output(
stream_mode,
@@ -3143,10 +3140,7 @@ class Pregel(
_state_mapper,
):
yield o
loop.after_tick()
# wait for checkpoint
if durability_ == "sync":
await cast(asyncio.Future, loop._put_checkpoint_fut)
await loop.aafter_tick()
finally:
# ensure waiter doesn't remain pending on cancel/shutdown
if _cleanup_waiter is not None:
+108 -11
View File
@@ -54,6 +54,13 @@ from langgraph.pregel import (
NodeBuilder,
Pregel,
)
from langgraph.pregel._checkpoint_writer import (
CHECKPOINT_BACKLOG_ENV_VAR,
DEFAULT_CHECKPOINT_BACKLOG,
AsyncCheckpointWriter,
SyncCheckpointWriter,
resolve_checkpoint_backlog,
)
from langgraph.pregel._loop import SyncPregelLoop
from langgraph.pregel._runner import PregelRunner
from langgraph.types import (
@@ -1329,20 +1336,22 @@ def test_imp_nested(
}
thread1 = {"configurable": {"thread_id": "1"}}
assert [*graph.stream([0, 1], thread1, durability=durability)] == [
{"submapper": "0"},
result = [*graph.stream([0, 1], thread1, durability=durability)]
# nested tasks run concurrently so output order is non-deterministic
assert sorted(result[:-1], key=lambda d: str(d)) == [
{"mapper": "00"},
{"submapper": "1"},
{"mapper": "11"},
{
"__interrupt__": (
Interrupt(
value="question",
id=AnyStr(),
),
)
},
{"submapper": "0"},
{"submapper": "1"},
]
assert result[-1] == {
"__interrupt__": (
Interrupt(
value="question",
id=AnyStr(),
),
)
}
assert graph.invoke(Command(resume="answer"), thread1, durability=durability) == [
"00answera",
@@ -3724,6 +3733,7 @@ def test_repeat_condition(snapshot: SnapshotAssertion) -> None:
"end": END,
},
)
workflow.add_conditional_edges(
"Chart Generator",
router,
@@ -3747,6 +3757,93 @@ def test_repeat_condition(snapshot: SnapshotAssertion) -> None:
assert app.get_graph().draw_mermaid(with_styles=False) == snapshot
def test_sync_durability_applies_checkpoint_backpressure() -> None:
first_put_started = threading.Event()
release_first_put = threading.Event()
put_calls = 0
visited: list[int] = []
result: dict[str, Any] = {}
error: dict[str, BaseException] = {}
class SlowFirstPutCheckpointer(InMemorySaver):
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: Any,
) -> RunnableConfig:
nonlocal put_calls
put_calls += 1
if put_calls == 1:
first_put_started.set()
release_first_put.wait()
return super().put(config, checkpoint, metadata, new_versions)
class State(TypedDict):
counter: int
def increment(state: State) -> State:
visited.append(state["counter"])
return {"counter": state["counter"] + 1}
def should_continue(state: State) -> str:
return "loop" if state["counter"] < 4 else "done"
builder = StateGraph(State)
builder.add_node("increment", increment)
builder.add_edge(START, "increment")
builder.add_conditional_edges(
"increment", should_continue, {"loop": "increment", "done": END}
)
graph = builder.compile(checkpointer=SlowFirstPutCheckpointer())
def invoke() -> None:
try:
result["value"] = graph.invoke(
{"counter": 0},
{"configurable": {"thread_id": "1"}},
durability="async",
)
except BaseException as exc:
error["value"] = exc
thread = threading.Thread(target=invoke)
thread.start()
assert first_put_started.wait(timeout=1)
time.sleep(0.05)
assert thread.is_alive()
assert len(visited) <= DEFAULT_CHECKPOINT_BACKLOG + 1
release_first_put.set()
thread.join(timeout=1)
assert not thread.is_alive()
assert "value" not in error
assert result["value"] == {"counter": 4}
def test_checkpoint_backlog_uses_env_override(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv(CHECKPOINT_BACKLOG_ENV_VAR, "7")
assert resolve_checkpoint_backlog() == 7
assert SyncCheckpointWriter(lambda *_args: None).queue.maxsize == 7
assert AsyncCheckpointWriter(lambda *_args: None).queue.maxsize == 7
def test_checkpoint_backlog_invalid_env_uses_default(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv(CHECKPOINT_BACKLOG_ENV_VAR, "not-an-int")
assert resolve_checkpoint_backlog() == DEFAULT_CHECKPOINT_BACKLOG
monkeypatch.setenv(CHECKPOINT_BACKLOG_ENV_VAR, "0")
assert resolve_checkpoint_backlog() == DEFAULT_CHECKPOINT_BACKLOG
def test_checkpoint_metadata(sync_checkpointer: BaseCheckpointSaver) -> None:
"""This test verifies that a run's configurable fields are merged with the
previous checkpoint config for each step in the run.
+56
View File
@@ -53,6 +53,7 @@ from langgraph.func import entrypoint, task
from langgraph.graph import END, START, StateGraph
from langgraph.graph.message import MessagesState, add_messages
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.pregel._checkpoint_writer import DEFAULT_CHECKPOINT_BACKLOG
from langgraph.pregel._loop import AsyncPregelLoop
from langgraph.pregel._runner import PregelRunner
from langgraph.types import (
@@ -484,6 +485,61 @@ async def test_checkpoint_put_after_cancellation_stream_events_anext() -> None:
assert False, "Task should be cancelled"
async def test_async_durability_applies_checkpoint_backpressure() -> None:
first_put_started = asyncio.Event()
release_first_put = asyncio.Event()
put_calls = 0
visited: list[int] = []
class SlowFirstPutCheckpointer(InMemorySaver):
async def aput(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
nonlocal put_calls
put_calls += 1
if put_calls == 1:
first_put_started.set()
await release_first_put.wait()
return await super().aput(config, checkpoint, metadata, new_versions)
class State(TypedDict):
counter: int
def increment(state: State) -> State:
visited.append(state["counter"])
return {"counter": state["counter"] + 1}
def should_continue(state: State) -> str:
return "loop" if state["counter"] < 4 else "done"
builder = StateGraph(State)
builder.add_node("increment", increment)
builder.add_edge(START, "increment")
builder.add_conditional_edges(
"increment", should_continue, {"loop": "increment", "done": END}
)
graph = builder.compile(checkpointer=SlowFirstPutCheckpointer())
task = asyncio.create_task(
graph.ainvoke(
{"counter": 0}, {"configurable": {"thread_id": "1"}}, durability="async"
)
)
await first_put_started.wait()
await asyncio.sleep(0.05)
assert not task.done()
assert len(visited) <= DEFAULT_CHECKPOINT_BACKLOG + 1
release_first_put.set()
assert await task == {"counter": 4}
async def test_node_cancellation_on_external_cancel() -> None:
inner_task_cancelled = False