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 |
@@ -26,7 +26,6 @@ from langgraph._internal._constants import (
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
NS_END,
|
||||
NS_SEP,
|
||||
NULL_TASK_ID,
|
||||
PUSH,
|
||||
SNAPSHOT_BUMPS,
|
||||
)
|
||||
@@ -65,14 +64,10 @@ def exit_delta_task_id(step: int, task_id: str) -> str:
|
||||
|
||||
Embeds the superstep in the first UUID group so `ORDER BY task_id, idx`
|
||||
preserves chronological order while remaining a valid RFC UUID (required by
|
||||
Postgres `checkpoint_writes.task_id uuid` columns). Never `NULL_TASK_ID`:
|
||||
readers apply writes under it as the anchor checkpoint's own pending writes.
|
||||
Postgres `checkpoint_writes.task_id uuid` columns).
|
||||
"""
|
||||
parts = str(uuid.UUID(task_id)).split("-")
|
||||
synthetic = f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
|
||||
if synthetic == NULL_TASK_ID:
|
||||
return f"{step:08d}-0000-0000-0000-000000000001"
|
||||
return synthetic
|
||||
return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
|
||||
|
||||
|
||||
def exit_delta_late_task_id(step: int, task_id: str) -> str:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -17,7 +17,6 @@ from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.checkpoint.serde.types import _DeltaSnapshot
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph._internal._constants import NULL_TASK_ID
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import START, StateGraph
|
||||
from langgraph.graph.message import _messages_delta_reducer
|
||||
@@ -39,7 +38,6 @@ def test_exit_delta_task_id_is_valid_uuid_and_ordered() -> None:
|
||||
assert id1.split("-")[0] == "00000001"
|
||||
assert id7.split("-")[0] == "00000007"
|
||||
assert id1.endswith("-0270-bf16-1ef8-fb321bef9f3d")
|
||||
assert exit_delta_task_id(0, NULL_TASK_ID) != NULL_TASK_ID
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
uuid.UUID(f"00000001-{tid}")
|
||||
@@ -465,43 +463,6 @@ def test_resume_with_a_command_update_replays_its_write_once(
|
||||
assert state.values["log"] == state.values["plain"] == ["in", "cmd", "ask", "done"]
|
||||
|
||||
|
||||
def test_command_update_on_an_input_checkpoint_matches_a_plain_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver, durability: Durability
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("node", lambda state: _both("node"))
|
||||
builder.add_edge(START, "node")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
graph.update_state(config, _both("in"), as_node="__input__")
|
||||
|
||||
graph.invoke(Command(update=_both("cmd")), config, durability=durability)
|
||||
|
||||
history = list(graph.get_state_history(config))
|
||||
assert [s.values.get("log", []) for s in history] == [
|
||||
s.values.get("plain", []) for s in history
|
||||
]
|
||||
replayed = graph.invoke(None, history[-1].config, durability=durability)
|
||||
assert replayed["log"] == replayed["plain"]
|
||||
|
||||
|
||||
def test_exit_command_update_on_a_new_thread_matches_a_plain_channel(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
builder = StateGraph(_ResumeState)
|
||||
builder.add_node("node", lambda state: _both("node"))
|
||||
builder.add_edge(START, "node")
|
||||
graph = builder.compile(checkpointer=sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "t"}}
|
||||
|
||||
graph.invoke(Command(update=_both("cmd")), config, durability="exit")
|
||||
|
||||
history = list(graph.get_state_history(config))
|
||||
assert [s.values.get("log", []) for s in history] == [
|
||||
s.values.get("plain", []) for s in history
|
||||
]
|
||||
|
||||
|
||||
class _FlagState(_ResumeState, total=False):
|
||||
extra: Annotated[list, DeltaChannel(_append)]
|
||||
flag: bool
|
||||
|
||||
@@ -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
|
||||
@@ -1972,7 +1972,7 @@ class AsyncThreadStream:
|
||||
# Mark that we have observed an active run so thread.output
|
||||
# knows a run exists (handles reattach without run.start).
|
||||
self._run_seen = True
|
||||
elif _is_root_terminal_lifecycle(event):
|
||||
elif phase in ("completed", "failed"):
|
||||
# Why: interrupts describe current-run state. Clear on terminal
|
||||
# lifecycle so a subsequent run.respond() can't fire against a
|
||||
# stale prior-run interrupt_id. Acquire `_interrupts_lock` so
|
||||
|
||||
@@ -35,11 +35,7 @@ from langgraph_sdk.stream.decoders import (
|
||||
validate_interleave_channels,
|
||||
)
|
||||
from langgraph_sdk.stream.subscription import compute_union_filter, infer_channel
|
||||
from langgraph_sdk.stream.sync_controller import (
|
||||
SyncStreamController,
|
||||
_is_root_terminal_lifecycle,
|
||||
_SyncSubscription,
|
||||
)
|
||||
from langgraph_sdk.stream.sync_controller import SyncStreamController, _SyncSubscription
|
||||
from langgraph_sdk.stream.transport import (
|
||||
SyncEventStreamHandle,
|
||||
SyncProtocolSseTransport,
|
||||
@@ -1618,7 +1614,7 @@ class SyncThreadStream:
|
||||
phase = data.get("event") if isinstance(data, dict) else None
|
||||
if phase in ("started", "running"):
|
||||
self._run_seen = True
|
||||
elif _is_root_terminal_lifecycle(event):
|
||||
elif phase in ("completed", "failed"):
|
||||
self.interrupted = False
|
||||
self.interrupts = []
|
||||
run_done = self._run_done
|
||||
|
||||
@@ -15,7 +15,6 @@ from langgraph_sdk.stream.transport import EventStreamHandle, ProtocolSseTranspo
|
||||
from streaming._events import (
|
||||
input_requested_event,
|
||||
lifecycle_completed_event,
|
||||
lifecycle_errored_event,
|
||||
lifecycle_event,
|
||||
)
|
||||
from streaming._fake_server import FakeServer, _StreamScript
|
||||
@@ -116,25 +115,6 @@ async def test_terminal_lifecycle_clears_interrupts():
|
||||
assert thread.interrupts == []
|
||||
|
||||
|
||||
async def test_subgraph_completed_event_does_not_end_run():
|
||||
fake = FakeServer()
|
||||
fake.script(
|
||||
[
|
||||
lifecycle_completed_event(seq=0, namespace=["child:1"]),
|
||||
lifecycle_errored_event(seq=1, error="root failed"),
|
||||
]
|
||||
)
|
||||
asgi = httpx.ASGITransport(app=fake.app)
|
||||
async with httpx.AsyncClient(transport=asgi, base_url="http://test") as raw:
|
||||
threads = ThreadsClient(HttpClient(raw))
|
||||
async with threads.stream(thread_id="t-1", assistant_id="agent") as thread:
|
||||
run_done = thread._run_done
|
||||
assert run_done is not None
|
||||
terminal = await asyncio.wait_for(run_done, timeout=2.0)
|
||||
assert terminal.status == "errored", "a subgraph's completed event ended the run"
|
||||
assert "root failed" in str(terminal.error)
|
||||
|
||||
|
||||
async def test_lifecycle_error_captured_for_output():
|
||||
"""Lifecycle error terminal state is captured in _run_done with error set."""
|
||||
fake = FakeServer()
|
||||
|
||||
@@ -27,7 +27,6 @@ from streaming._events import (
|
||||
checkpoints_event,
|
||||
custom_event,
|
||||
lifecycle_completed_event,
|
||||
lifecycle_errored_event,
|
||||
lifecycle_event,
|
||||
lifecycle_started_event,
|
||||
message_finish_event,
|
||||
@@ -476,23 +475,6 @@ def test_sync_lifecycle_watcher_reconnects_with_since_after_transport_drop():
|
||||
assert fake.stream_request_bodies[1]["since"] == 1
|
||||
|
||||
|
||||
def test_sync_subgraph_completed_event_does_not_end_run():
|
||||
fake = SyncFakeServer()
|
||||
fake.script(
|
||||
[
|
||||
lifecycle_completed_event(seq=1, namespace=["child:1"]),
|
||||
lifecycle_errored_event(seq=2, error="root failed"),
|
||||
]
|
||||
)
|
||||
with httpx.Client(transport=fake.transport, base_url="http://test") as raw:
|
||||
threads = SyncThreadsClient(SyncHttpClient(raw))
|
||||
with threads.stream(thread_id="existing", assistant_id="agent") as thread:
|
||||
terminal = thread._wait_for_run_done()
|
||||
|
||||
assert terminal.status == "errored", "a subgraph's completed event ended the run"
|
||||
assert "root failed" in str(terminal.error)
|
||||
|
||||
|
||||
def test_sync_threads_stream_accepts_websocket_transport_option():
|
||||
with httpx.Client(base_url="http://test") as raw:
|
||||
threads = SyncThreadsClient(SyncHttpClient(raw))
|
||||
|
||||
Reference in New Issue
Block a user