chore: use bounded queue for checkpointing

This commit is contained in:
Will Fu-Hinthorn
2026-04-09 17:33:50 -07:00
parent bde0c47cd9
commit 685a755baf
5 changed files with 476 additions and 97 deletions
@@ -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:
+95
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 (
@@ -3726,6 +3733,7 @@ def test_repeat_condition(snapshot: SnapshotAssertion) -> None:
"end": END,
},
)
workflow.add_conditional_edges(
"Chart Generator",
router,
@@ -3749,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