mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-08 17:35:17 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b791c1f1d5 |
@@ -191,6 +191,7 @@ class PregelLoop:
|
||||
Callable[
|
||||
[
|
||||
concurrent.futures.Future | None,
|
||||
Sequence[Any],
|
||||
RunnableConfig,
|
||||
Checkpoint,
|
||||
str,
|
||||
@@ -204,11 +205,13 @@ class PregelLoop:
|
||||
submit: Submit
|
||||
channels: Mapping[str, BaseChannel]
|
||||
# Futures from `checkpointer.put_writes` calls that produced delta-channel
|
||||
# writes. `_checkpointer_put_after_previous` drains this list (swap to a
|
||||
# local `futs` then reset to `[]` and wait/gather) before putting the
|
||||
# next checkpoint, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. Initialised to `[]` in both sync and async
|
||||
# `__enter__`; stays `None` only when no checkpointer.
|
||||
# writes. `_put_checkpoint` hands this list to the save it submits, which
|
||||
# waits for them first, so a checkpoint never becomes durable before the
|
||||
# writes that produced it. If a write or the previous save failed, the
|
||||
# save fails too: a DeltaChannel is rebuilt from its writes along the
|
||||
# parent chain, so a checkpoint saved past either gap reads back short
|
||||
# for good. Initialised to `[]` in both sync and async `__enter__`;
|
||||
# stays `None` only when no checkpointer.
|
||||
_delta_write_futs: list[Any] | None = None
|
||||
|
||||
# Same pattern as `_delta_write_futs` but for error-handler writes.
|
||||
@@ -1299,12 +1302,17 @@ class PregelLoop:
|
||||
)
|
||||
self.checkpoint_previous_versions = channel_versions
|
||||
|
||||
# Take this checkpoint's writes now: saves run in the background
|
||||
# and can start out of order, so a save that took them itself
|
||||
# could get another checkpoint's writes.
|
||||
delta_write_futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
# 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),
|
||||
delta_write_futs,
|
||||
self.checkpoint_config,
|
||||
copy_checkpoint(self.checkpoint),
|
||||
self.checkpoint_metadata,
|
||||
@@ -1374,6 +1382,7 @@ class PregelLoop:
|
||||
self._put_checkpoint_fut = self.submit(
|
||||
self._checkpointer_put_after_previous,
|
||||
getattr(self, "_put_checkpoint_fut", None),
|
||||
(),
|
||||
stub_put_config,
|
||||
stub_cp,
|
||||
{"step": -2},
|
||||
@@ -1641,21 +1650,19 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
def _checkpointer_put_after_previous(
|
||||
self,
|
||||
prev: concurrent.futures.Future | None,
|
||||
delta_write_futs: Sequence[concurrent.futures.Future],
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
concurrent.futures.wait(futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
finally:
|
||||
cast(BaseCheckpointSaver, self.checkpointer).put(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
for fut in delta_write_futs:
|
||||
fut.result()
|
||||
if prev is not None:
|
||||
prev.result()
|
||||
cast(BaseCheckpointSaver, self.checkpointer).put(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
|
||||
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
if self.cache is None:
|
||||
@@ -1896,23 +1903,19 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
async def _checkpointer_put_after_previous(
|
||||
self,
|
||||
prev: asyncio.Task | None,
|
||||
delta_write_futs: Sequence[asyncio.Future],
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
new_versions: ChannelVersions,
|
||||
) -> RunnableConfig:
|
||||
# Drain DeltaChannel write futures before committing the checkpoint so
|
||||
# ancestor walks never see a checkpoint without its backing writes.
|
||||
if self._delta_write_futs:
|
||||
futs, self._delta_write_futs = self._delta_write_futs, []
|
||||
await asyncio.gather(*futs)
|
||||
try:
|
||||
if prev is not None:
|
||||
await prev
|
||||
finally:
|
||||
await cast(BaseCheckpointSaver, self.checkpointer).aput(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
if delta_write_futs:
|
||||
await asyncio.gather(*delta_write_futs)
|
||||
if prev is not None:
|
||||
await prev
|
||||
await cast(BaseCheckpointSaver, self.checkpointer).aput(
|
||||
config, checkpoint, metadata, new_versions
|
||||
)
|
||||
|
||||
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
if self.cache is None:
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""A checkpoint must never be saved without the `DeltaChannel` writes it reads."""
|
||||
|
||||
import operator
|
||||
import threading
|
||||
from typing import Annotated, Any
|
||||
|
||||
import pytest
|
||||
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.types import Durability
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
INPUT = {"log": [], "plain": []}
|
||||
FINAL = {"log": ["a", "b", "c"], "plain": ["a", "b", "c"]}
|
||||
|
||||
# Exit mode saves nothing before the failed write, so its retry starts over.
|
||||
RETRIES = [
|
||||
pytest.param("sync", None, id="sync"),
|
||||
pytest.param("async", None, id="async"),
|
||||
pytest.param("exit", INPUT, id="exit"),
|
||||
]
|
||||
|
||||
|
||||
def _append(current: list, writes: list) -> list:
|
||||
return [*current, *(item for write in writes for item in write)]
|
||||
|
||||
|
||||
class _State(TypedDict):
|
||||
log: Annotated[list, DeltaChannel(_append)]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
|
||||
class _FailsTheWriteOfBOnce(InMemorySaver):
|
||||
failed = False
|
||||
|
||||
def _fail_once(self, writes: Any) -> None:
|
||||
if not self.failed and ("log", ["b"]) in writes:
|
||||
self.failed = True
|
||||
raise ConnectionError("b's write was not saved")
|
||||
|
||||
def put_writes(
|
||||
self, config: Any, writes: Any, task_id: str, task_path: str = ""
|
||||
) -> None:
|
||||
self._fail_once(writes)
|
||||
super().put_writes(config, writes, task_id, task_path)
|
||||
|
||||
async def aput_writes(
|
||||
self, config: Any, writes: Any, task_id: str, task_path: str = ""
|
||||
) -> None:
|
||||
self._fail_once(writes)
|
||||
await super().aput_writes(config, writes, task_id, task_path)
|
||||
|
||||
|
||||
def _a_then_b_then_c(saver: InMemorySaver) -> Any:
|
||||
builder = StateGraph(_State)
|
||||
for name in "abc":
|
||||
builder.add_node(
|
||||
name, lambda state, name=name: {"log": [name], "plain": [name]}
|
||||
)
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("b", "c")
|
||||
return builder.compile(checkpointer=saver)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
|
||||
def test_a_failed_delta_write_is_rerun_not_lost(
|
||||
durability: Durability, retry_input: dict | None
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
graph.invoke(INPUT, config, durability=durability)
|
||||
for state in graph.get_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
graph.invoke(retry_input, config, durability=durability)
|
||||
assert graph.get_state(config).values == FINAL
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("durability", "retry_input"), RETRIES)
|
||||
async def test_a_failed_delta_write_is_rerun_not_lost_async(
|
||||
durability: Durability, retry_input: dict | None
|
||||
) -> None:
|
||||
graph = _a_then_b_then_c(_FailsTheWriteOfBOnce())
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
with pytest.raises(ConnectionError):
|
||||
await graph.ainvoke(INPUT, config, durability=durability)
|
||||
async for state in graph.aget_state_history(config):
|
||||
assert state.values.get("log", []) == state.values.get("plain", [])
|
||||
|
||||
await graph.ainvoke(retry_input, config, durability=durability)
|
||||
assert (await graph.aget_state(config)).values == FINAL
|
||||
|
||||
|
||||
def test_a_delta_graph_finishes_on_a_single_background_thread() -> None:
|
||||
graph = _a_then_b_then_c(InMemorySaver())
|
||||
config = {"configurable": {"thread_id": "t"}, "max_concurrency": 1}
|
||||
result: dict = {}
|
||||
run = threading.Thread(
|
||||
target=lambda: result.update(graph.invoke(INPUT, config, durability="async")),
|
||||
daemon=True,
|
||||
)
|
||||
|
||||
run.start()
|
||||
run.join(timeout=10)
|
||||
|
||||
assert not run.is_alive(), "invoke hung"
|
||||
assert result == FINAL
|
||||
Reference in New Issue
Block a user