mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
chore: use bounded queue for checkpointing
This commit is contained in:
@@ -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)
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user