Compare commits

..
Author SHA1 Message Date
Elior Nataf LackritzandGitHub 40a2e6d845 fix(langgraph): never store exit-mode delta writes under the null task id (#9229)
In exit mode, a run stores its DeltaChannel writes on the checkpoint it
started from, under synthetic task ids that put the superstep first. For
a `Command`'s writes in a run whose first superstep is 0, that id came
out as `NULL_TASK_ID` itself. Readers take writes under `NULL_TASK_ID`
for the checkpoint's own pending writes, so after
`invoke(Command(update=...), durability="exit")`:

- on a new thread, `get_state` on the empty first checkpoint showed the
update in the DeltaChannel but not in a plain channel;
- on a thread whose first checkpoint came from `update_state(...,
as_node="__input__")`, `get_state` on that checkpoint showed the update
the same way, and a replay or fork from it applied the update to the
DeltaChannel only.

`exit_delta_task_id` now never returns `NULL_TASK_ID`. Threads already
saved this way stay as they are.

## Tests

- `test_command_update_on_an_input_checkpoint_matches_a_plain_channel`
compares every checkpoint in the history, and a replay from the input
checkpoint, with a plain channel, on every checkpointer and durability.
The exit cases fail without the fix.
- `test_exit_command_update_on_a_new_thread_matches_a_plain_channel`
compares every checkpoint in a new thread's history with a plain
channel, on every checkpointer, and fails without the fix. It runs in
exit mode only: in sync and async the two channels already differ on
`main`, because of a separate bug that applies a new thread's `Command`
update twice.

`make format`, `make lint` and the langgraph suite pass.
2026-10-07 19:17:02 -04:00
Elior Nataf LackritzandGitHub a0053bb616 fix(sdk-py): end a thread stream's run only on a root lifecycle event (#9228)
The thread stream's lifecycle watcher ended the run on any `completed`
or `failed` lifecycle event, including the one a subgraph sends when it
finishes. So `thread.output` could read the thread state before the
parent stored the step that ran the subgraph, and a run that failed
after a subgraph completed was reported as completed. The async and sync
watchers now end the run only on a root event, with the
`_is_root_terminal_lifecycle` check the fanout already uses. The JS SDK
checks the root namespace here too.

This is the `sdk-py integration` flake where the final `items` comes
back as `['streamed', 'tool', 'asked']` without `'sub'`: the example
graph's last node runs a subgraph.

## Tests

New async and sync tests send a subgraph `completed` and then a root
`failed`. Without the fix the run ends as completed. `make lint` and
`make test` in `libs/sdk-py` pass.
2026-10-07 19:03:46 -04:00
8 changed files with 118 additions and 150 deletions
@@ -26,6 +26,7 @@ from langgraph._internal._constants import (
CONFIG_KEY_CHECKPOINT_ID,
NS_END,
NS_SEP,
NULL_TASK_ID,
PUSH,
SNAPSHOT_BUMPS,
)
@@ -64,10 +65,14 @@ 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).
Postgres `checkpoint_writes.task_id uuid` columns). Never `NULL_TASK_ID`:
readers apply writes under it as the anchor checkpoint's own pending writes.
"""
parts = str(uuid.UUID(task_id)).split("-")
return f"{step:08d}-{parts[1]}-{parts[2]}-{parts[3]}-{parts[4]}"
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
def exit_delta_late_task_id(step: int, task_id: str) -> str:
+27 -30
View File
@@ -191,7 +191,6 @@ class PregelLoop:
Callable[
[
concurrent.futures.Future | None,
Sequence[Any],
RunnableConfig,
Checkpoint,
str,
@@ -205,13 +204,11 @@ class PregelLoop:
submit: Submit
channels: Mapping[str, BaseChannel]
# Futures from `checkpointer.put_writes` calls that produced delta-channel
# 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.
# 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.
_delta_write_futs: list[Any] | None = None
# Same pattern as `_delta_write_futs` but for error-handler writes.
@@ -1302,17 +1299,12 @@ 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,
@@ -1382,7 +1374,6 @@ 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},
@@ -1650,19 +1641,21 @@ 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:
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
)
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
)
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
@@ -1903,19 +1896,23 @@ 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:
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
)
# 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
)
async def amatch_cached_writes(self) -> Sequence[PregelExecutableTask]:
if self.cache is None:
@@ -17,6 +17,7 @@ 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
@@ -38,6 +39,7 @@ 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}")
@@ -463,6 +465,43 @@ 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
@@ -1,115 +0,0 @@
"""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
+1 -1
View File
@@ -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 phase in ("completed", "failed"):
elif _is_root_terminal_lifecycle(event):
# 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
+6 -2
View File
@@ -35,7 +35,11 @@ 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, _SyncSubscription
from langgraph_sdk.stream.sync_controller import (
SyncStreamController,
_is_root_terminal_lifecycle,
_SyncSubscription,
)
from langgraph_sdk.stream.transport import (
SyncEventStreamHandle,
SyncProtocolSseTransport,
@@ -1614,7 +1618,7 @@ class SyncThreadStream:
phase = data.get("event") if isinstance(data, dict) else None
if phase in ("started", "running"):
self._run_seen = True
elif phase in ("completed", "failed"):
elif _is_root_terminal_lifecycle(event):
self.interrupted = False
self.interrupts = []
run_done = self._run_done
@@ -15,6 +15,7 @@ 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
@@ -115,6 +116,25 @@ 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,6 +27,7 @@ from streaming._events import (
checkpoints_event,
custom_event,
lifecycle_completed_event,
lifecycle_errored_event,
lifecycle_event,
lifecycle_started_event,
message_finish_event,
@@ -475,6 +476,23 @@ 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))