mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-08 17:35:17 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
40a2e6d845 | ||
|
|
a0053bb616 |
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user