mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-05 16:05:09 +02:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4826448e1a | ||
|
|
8ea279b8ad | ||
|
|
eb2f09e321 | ||
|
|
26d279a0ac | ||
|
|
e850b21d08 | ||
|
|
6a92b7ff3c |
@@ -41,6 +41,8 @@ CONFIG_KEY_CACHE = sys.intern("__pregel_cache")
|
||||
# holds a `BaseCache` made available to subgraphs
|
||||
CONFIG_KEY_RESUMING = sys.intern("__pregel_resuming")
|
||||
# holds a boolean indicating if subgraphs should resume from a previous checkpoint
|
||||
CONFIG_KEY_REPLAYING = sys.intern("__pregel_replaying")
|
||||
# holds a boolean indicating if subgraphs should replay (re-run tasks, drop cached RESUME writes)
|
||||
CONFIG_KEY_TASK_ID = sys.intern("__pregel_task_id")
|
||||
# holds the task ID for the current task
|
||||
CONFIG_KEY_THREAD_ID = sys.intern("thread_id")
|
||||
@@ -98,6 +100,7 @@ RESERVED = {
|
||||
CONFIG_KEY_STREAM,
|
||||
CONFIG_KEY_CHECKPOINT_MAP,
|
||||
CONFIG_KEY_RESUMING,
|
||||
CONFIG_KEY_REPLAYING,
|
||||
CONFIG_KEY_TASK_ID,
|
||||
CONFIG_KEY_CHECKPOINT_MAP,
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
|
||||
@@ -42,6 +42,7 @@ from langgraph._internal._constants import (
|
||||
CONFIG_KEY_CHECKPOINT_ID,
|
||||
CONFIG_KEY_CHECKPOINT_MAP,
|
||||
CONFIG_KEY_CHECKPOINT_NS,
|
||||
CONFIG_KEY_REPLAYING,
|
||||
CONFIG_KEY_RESUME_MAP,
|
||||
CONFIG_KEY_RESUMING,
|
||||
CONFIG_KEY_SCRATCHPAD,
|
||||
@@ -152,7 +153,7 @@ class PregelLoop:
|
||||
input_keys: str | Sequence[str]
|
||||
output_keys: str | Sequence[str]
|
||||
stream_keys: str | Sequence[str]
|
||||
skip_done_tasks: bool
|
||||
is_replaying: bool
|
||||
is_nested: bool
|
||||
manager: None | AsyncParentRunManager | ParentRunManager
|
||||
interrupt_after: All | Sequence[str]
|
||||
@@ -244,7 +245,9 @@ class PregelLoop:
|
||||
self.interrupt_before = interrupt_before
|
||||
self.manager = manager
|
||||
self.is_nested = CONFIG_KEY_TASK_ID in self.config.get(CONF, {})
|
||||
self.skip_done_tasks = CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
|
||||
self.is_replaying = CONFIG_KEY_CHECKPOINT_ID in config[CONF] or config[
|
||||
CONF
|
||||
].get(CONFIG_KEY_REPLAYING, False)
|
||||
self._migrate_checkpoint = migrate_checkpoint
|
||||
self.trigger_to_nodes = trigger_to_nodes
|
||||
self.retry_policy = retry_policy
|
||||
@@ -451,7 +454,7 @@ class PregelLoop:
|
||||
# save the new task
|
||||
self.tasks[pushed.id] = pushed
|
||||
# match any pending writes to the new task
|
||||
if self.skip_done_tasks:
|
||||
if not self.is_replaying:
|
||||
self._match_writes({pushed.id: pushed})
|
||||
# return the new task, to be started if not run before
|
||||
return pushed
|
||||
@@ -515,7 +518,7 @@ class PregelLoop:
|
||||
return False
|
||||
|
||||
# if there are pending writes from a previous loop, apply them
|
||||
if self.skip_done_tasks and self.checkpoint_pending_writes:
|
||||
if not self.is_replaying and self.checkpoint_pending_writes:
|
||||
self._match_writes(self.tasks)
|
||||
|
||||
# before execution, check if we should interrupt
|
||||
@@ -557,8 +560,8 @@ class PregelLoop:
|
||||
)
|
||||
# clear pending writes
|
||||
self.checkpoint_pending_writes.clear()
|
||||
# "not skip_done_tasks" only applies to first tick after resuming
|
||||
self.skip_done_tasks = True
|
||||
# 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
|
||||
@@ -567,8 +570,9 @@ class PregelLoop:
|
||||
):
|
||||
self.status = "interrupt_after"
|
||||
raise GraphInterrupt()
|
||||
# unset resuming flag
|
||||
# unset resuming/replaying flags
|
||||
self.config[CONF].pop(CONFIG_KEY_RESUMING, None)
|
||||
self.config[CONF].pop(CONFIG_KEY_REPLAYING, None)
|
||||
|
||||
def match_cached_writes(self) -> Sequence[PregelExecutableTask]:
|
||||
raise NotImplementedError
|
||||
@@ -641,6 +645,20 @@ class PregelLoop:
|
||||
configurable.get(CONFIG_KEY_RESUMING, input_signals_resume)
|
||||
)
|
||||
|
||||
# When replaying from a specific checkpoint, drop cached RESUME
|
||||
# writes so that interrupt() calls re-fire instead of returning
|
||||
# stale values. But if a resume value is being provided (e.g.
|
||||
# Command(resume=...) or CONFIG_KEY_RESUMING), keep them —
|
||||
# multi-interrupt scenarios need previously resolved values preserved.
|
||||
if self.is_replaying:
|
||||
is_resume_with_value = (
|
||||
isinstance(self.input, Command) and self.input.resume is not None
|
||||
) or configurable.get(CONFIG_KEY_RESUMING, False)
|
||||
if not is_resume_with_value:
|
||||
self.checkpoint_pending_writes = [
|
||||
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
||||
]
|
||||
|
||||
# map command to writes
|
||||
if isinstance(self.input, Command):
|
||||
if (resume := self.input.resume) is not None:
|
||||
@@ -729,18 +747,14 @@ class PregelLoop:
|
||||
self._put_checkpoint({"source": "input"})
|
||||
elif CONFIG_KEY_RESUMING not in configurable:
|
||||
raise EmptyInputError(f"Received no input for {input_keys}")
|
||||
# Propagate resuming flag to subgraphs (only the outer graph does this).
|
||||
# Propagate resuming and replaying flags to subgraphs.
|
||||
if not self.is_nested:
|
||||
has_resume_value = (
|
||||
isinstance(self.input, Command) and self.input.resume is not None
|
||||
)
|
||||
# When forking (skip_done_tasks=False, i.e. specific checkpoint_id),
|
||||
# subgraphs should NOT resume — they start fresh.
|
||||
# When genuinely resuming from latest, subgraphs should also resume.
|
||||
is_fork = not self.skip_done_tasks
|
||||
subgraph_should_resume = has_resume_value or (is_resuming and not is_fork)
|
||||
self.config = patch_configurable(
|
||||
self.config, {CONFIG_KEY_RESUMING: subgraph_should_resume}
|
||||
self.config,
|
||||
{
|
||||
CONFIG_KEY_RESUMING: is_resuming,
|
||||
CONFIG_KEY_REPLAYING: self.is_replaying,
|
||||
},
|
||||
)
|
||||
# set flag
|
||||
self.status = "pending"
|
||||
@@ -1092,6 +1106,59 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
},
|
||||
)
|
||||
|
||||
def _get_parent_checkpoint_id(self) -> str | None:
|
||||
"""Get the parent checkpoint_id to use as an upper bound for finding
|
||||
the subgraph's checkpoint. For forks, we need the original parent
|
||||
checkpoint (not the fork), so we look up the parent checkpoint's
|
||||
parent_config."""
|
||||
checkpoint_map = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
|
||||
parent_ns = NS_SEP.join(self.checkpoint_ns[:-1]) if self.checkpoint_ns else ""
|
||||
parent_checkpoint_id = checkpoint_map.get(parent_ns)
|
||||
if not parent_checkpoint_id or not self.checkpointer:
|
||||
return None
|
||||
# Check if this is a fork (source=update) — if so, use the fork's
|
||||
# parent checkpoint_id instead, since the fork was created after
|
||||
# the subgraph's checkpoints from the original execution.
|
||||
parent_config: RunnableConfig = {
|
||||
**self.checkpoint_config,
|
||||
CONF: {
|
||||
**self.checkpoint_config.get(CONF, {}),
|
||||
CONFIG_KEY_CHECKPOINT_NS: parent_ns,
|
||||
CONFIG_KEY_CHECKPOINT_ID: parent_checkpoint_id,
|
||||
},
|
||||
}
|
||||
parent_saved = self.checkpointer.get_tuple(parent_config)
|
||||
if parent_saved and parent_saved.metadata.get("source") == "update":
|
||||
if parent_saved.parent_config:
|
||||
return parent_saved.parent_config[CONF].get(CONFIG_KEY_CHECKPOINT_ID)
|
||||
return parent_checkpoint_id
|
||||
|
||||
def _get_checkpoint_before_parent(self) -> CheckpointTuple | None:
|
||||
"""Find the subgraph checkpoint that was current at the parent's
|
||||
checkpoint time, using the parent checkpoint_id as an upper bound.
|
||||
|
||||
Returns a CheckpointTuple with the historical channel_values but fresh
|
||||
execution state (empty channel_versions/versions_seen) so the subgraph
|
||||
re-runs its nodes while retaining accumulated data.
|
||||
|
||||
Returns None to start fresh if no such checkpoint exists."""
|
||||
parent_checkpoint_id = self._get_parent_checkpoint_id()
|
||||
if parent_checkpoint_id and self.checkpointer:
|
||||
before_config: RunnableConfig = {
|
||||
CONF: {"checkpoint_id": parent_checkpoint_id}
|
||||
}
|
||||
for saved in self.checkpointer.list(
|
||||
self.checkpoint_config, before=before_config, limit=1
|
||||
):
|
||||
checkpoint = empty_checkpoint()
|
||||
checkpoint["channel_values"] = saved.checkpoint.get(
|
||||
"channel_values", {}
|
||||
)
|
||||
return CheckpointTuple(
|
||||
self.checkpoint_config, checkpoint, {"step": -2}, None, []
|
||||
)
|
||||
return None
|
||||
|
||||
# context manager
|
||||
|
||||
def __enter__(self) -> Self:
|
||||
@@ -1099,6 +1166,16 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
saved = self.checkpointer.get_tuple(self.checkpoint_config)
|
||||
else:
|
||||
saved = None
|
||||
# When replaying a subgraph that wasn't in the checkpoint map
|
||||
# (parent checkpoint predates this subgraph), start fresh.
|
||||
# For stateful subgraphs (checkpointer=True), find the checkpoint
|
||||
# that was current at the parent's checkpoint time.
|
||||
if (
|
||||
saved is not None
|
||||
and self.config[CONF].get(CONFIG_KEY_REPLAYING)
|
||||
and not self.checkpoint_config.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_ID)
|
||||
):
|
||||
saved = self._get_checkpoint_before_parent()
|
||||
if saved is None:
|
||||
saved = CheckpointTuple(
|
||||
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
|
||||
@@ -1123,20 +1200,6 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
|
||||
if saved.pending_writes is not None
|
||||
else []
|
||||
)
|
||||
# When replaying from a specific checkpoint (fork), drop cached
|
||||
# RESUME writes so that interrupt() calls re-fire instead of
|
||||
# returning stale values. But if the input directly carries a
|
||||
# resume value, keep them — multi-interrupt scenarios need
|
||||
# previously resolved RESUME values preserved.
|
||||
is_replaying = not self.skip_done_tasks
|
||||
has_resume_value = (
|
||||
self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING) is True
|
||||
) or (isinstance(self.input, Command) and self.input.resume is not None)
|
||||
if is_replaying and not has_resume_value:
|
||||
self.checkpoint_pending_writes = [
|
||||
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
||||
]
|
||||
|
||||
self.submit = self.stack.enter_context(BackgroundExecutor(self.config))
|
||||
self.channels, self.managed = channels_from_checkpoint(
|
||||
self.specs, self.checkpoint
|
||||
@@ -1284,6 +1347,48 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
},
|
||||
)
|
||||
|
||||
async def _aget_parent_checkpoint_id(self) -> str | None:
|
||||
checkpoint_map = self.config[CONF].get(CONFIG_KEY_CHECKPOINT_MAP, {})
|
||||
parent_ns = NS_SEP.join(self.checkpoint_ns[:-1]) if self.checkpoint_ns else ""
|
||||
parent_checkpoint_id = checkpoint_map.get(parent_ns)
|
||||
if not parent_checkpoint_id or not self.checkpointer:
|
||||
return None
|
||||
parent_config: RunnableConfig = {
|
||||
**self.checkpoint_config,
|
||||
CONF: {
|
||||
**self.checkpoint_config.get(CONF, {}),
|
||||
CONFIG_KEY_CHECKPOINT_NS: parent_ns,
|
||||
CONFIG_KEY_CHECKPOINT_ID: parent_checkpoint_id,
|
||||
},
|
||||
}
|
||||
parent_saved = await self.checkpointer.aget_tuple(parent_config)
|
||||
if parent_saved and parent_saved.metadata.get("source") == "update":
|
||||
if parent_saved.parent_config:
|
||||
return parent_saved.parent_config[CONF].get(CONFIG_KEY_CHECKPOINT_ID)
|
||||
return parent_checkpoint_id
|
||||
|
||||
async def _aget_checkpoint_before_parent(self) -> CheckpointTuple | None:
|
||||
parent_checkpoint_id = await self._aget_parent_checkpoint_id()
|
||||
if parent_checkpoint_id and self.checkpointer:
|
||||
before_config: RunnableConfig = {
|
||||
CONF: {"checkpoint_id": parent_checkpoint_id}
|
||||
}
|
||||
async for saved in self.checkpointer.alist(
|
||||
self.checkpoint_config, before=before_config, limit=1
|
||||
):
|
||||
checkpoint = empty_checkpoint()
|
||||
checkpoint["channel_values"] = saved.checkpoint.get(
|
||||
"channel_values", {}
|
||||
)
|
||||
return CheckpointTuple(
|
||||
self.checkpoint_config,
|
||||
checkpoint,
|
||||
{"step": -2},
|
||||
None,
|
||||
[],
|
||||
)
|
||||
return None
|
||||
|
||||
# context manager
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
@@ -1291,6 +1396,16 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
saved = await self.checkpointer.aget_tuple(self.checkpoint_config)
|
||||
else:
|
||||
saved = None
|
||||
# When replaying a subgraph that wasn't in the checkpoint map
|
||||
# (parent checkpoint predates this subgraph), start fresh.
|
||||
# For stateful subgraphs (checkpointer=True), find the checkpoint
|
||||
# that was current at the parent's checkpoint time.
|
||||
if (
|
||||
saved is not None
|
||||
and self.config[CONF].get(CONFIG_KEY_REPLAYING)
|
||||
and not self.checkpoint_config.get(CONF, {}).get(CONFIG_KEY_CHECKPOINT_ID)
|
||||
):
|
||||
saved = await self._aget_checkpoint_before_parent()
|
||||
if saved is None:
|
||||
saved = CheckpointTuple(
|
||||
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
|
||||
@@ -1315,20 +1430,6 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
|
||||
if saved.pending_writes is not None
|
||||
else []
|
||||
)
|
||||
# When replaying from a specific checkpoint (fork), drop cached
|
||||
# RESUME writes so that interrupt() calls re-fire instead of
|
||||
# returning stale values. But if the input directly carries a
|
||||
# resume value, keep them — multi-interrupt scenarios need
|
||||
# previously resolved RESUME values preserved.
|
||||
is_replaying = not self.skip_done_tasks
|
||||
has_resume_value = (
|
||||
self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING) is True
|
||||
) or (isinstance(self.input, Command) and self.input.resume is not None)
|
||||
if is_replaying and not has_resume_value:
|
||||
self.checkpoint_pending_writes = [
|
||||
w for w in self.checkpoint_pending_writes if w[1] != RESUME
|
||||
]
|
||||
|
||||
self.submit = await self.stack.enter_async_context(
|
||||
AsyncBackgroundExecutor(self.config)
|
||||
)
|
||||
|
||||
@@ -781,10 +781,15 @@ def test_subgraph_interrupt_replay_from_parent(
|
||||
history = list(graph.get_state_history(config))
|
||||
before_sub = [s for s in history if s.next == ("subgraph_node",)][-1]
|
||||
|
||||
# Replay — interrupt re-fires
|
||||
# Replay from before subgraph — subgraph starts fresh, interrupt re-fires
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, before_sub.config)
|
||||
assert "__interrupt__" in replay_result
|
||||
# Subgraph ran from scratch (step_a and ask_human called)
|
||||
assert "step_a" in called
|
||||
assert "ask_human" in called
|
||||
# step_b should NOT be called (interrupt stops execution)
|
||||
assert "step_b" not in called
|
||||
|
||||
|
||||
def test_subgraph_interrupt_fork_from_parent(
|
||||
@@ -930,6 +935,11 @@ def test_subgraph_interrupt_replay_from_interrupt_checkpoint(
|
||||
called.clear()
|
||||
replay_result = graph.invoke(None, interrupt_checkpoint.config)
|
||||
assert "__interrupt__" in replay_result
|
||||
# Subgraph starts fresh during replay — all nodes re-run from scratch.
|
||||
# step_a re-runs, ask_human re-fires interrupt, step_b not reached.
|
||||
assert "step_a" in called
|
||||
assert "ask_human" in called
|
||||
assert "step_b" not in called
|
||||
|
||||
|
||||
def test_subgraph_interrupt_fork_no_sub_checkpointer(
|
||||
@@ -1548,3 +1558,274 @@ def test_checkpoint_ns_accessible_in_subgraph(
|
||||
assert captured_config["checkpoint_ns"] is not None
|
||||
assert captured_config["checkpoint_ns"] != ""
|
||||
assert captured_config["thread_id"] == "1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Section 8: Stateful vs stateless subgraph state retention on replay
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_stateful_subgraph_retains_state_on_parent_replay(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""With checkpointer=True, the subgraph retains its accumulated state
|
||||
across parent invocations. Graph: parent_node -> sub_node.
|
||||
|
||||
Invoke the graph twice on the same thread. Each invocation triggers
|
||||
two interrupts (step_a, step_b). After both complete, replay from the
|
||||
checkpoint before sub_node in the 2nd invocation. The stateful subgraph
|
||||
should see accumulated state from the 1st invocation (a:Alice, b:30)
|
||||
but not the 2nd invocation's state (a:Bob, b:25)."""
|
||||
|
||||
observed_sub_input: list[tuple[str, dict]] = []
|
||||
|
||||
class SubState(TypedDict):
|
||||
value: Annotated[list[str], operator.add]
|
||||
|
||||
class ParentState(TypedDict):
|
||||
results: Annotated[list[str], operator.add]
|
||||
|
||||
def parent_node(state: ParentState) -> ParentState:
|
||||
return {"results": ["p"]}
|
||||
|
||||
def step_a(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_a", dict(state)))
|
||||
answer = interrupt("Question A")
|
||||
return {"value": [f"a:{answer}"]}
|
||||
|
||||
def step_b(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_b", dict(state)))
|
||||
answer = interrupt("Question B")
|
||||
return {"value": [f"b:{answer}"]}
|
||||
|
||||
sub = (
|
||||
StateGraph(SubState)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("step_b", step_b)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "step_b")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("parent_node", parent_node)
|
||||
.add_node("sub_node", sub)
|
||||
.add_edge(START, "parent_node")
|
||||
.add_edge("parent_node", "sub_node")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# 1st invocation: complete with 2 interrupts
|
||||
graph.invoke({"results": []}, config) # step_a interrupt
|
||||
graph.invoke(Command(resume="Alice"), config) # step_b interrupt
|
||||
graph.invoke(Command(resume="30"), config) # complete
|
||||
|
||||
# Verify 1st invocation: subgraph started fresh
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": []})
|
||||
assert step_b_entries[0] == ("step_b", {"value": ["a:Alice"]})
|
||||
|
||||
# 2nd invocation: complete with 2 interrupts
|
||||
observed_sub_input.clear()
|
||||
graph.invoke({"results": []}, config) # step_a interrupt
|
||||
graph.invoke(Command(resume="Bob"), config) # step_b interrupt
|
||||
graph.invoke(Command(resume="25"), config) # complete
|
||||
|
||||
# Verify 2nd invocation: subgraph retained state from 1st invocation
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": ["a:Alice", "b:30"]})
|
||||
assert step_b_entries[0] == (
|
||||
"step_b",
|
||||
{"value": ["a:Alice", "b:30", "a:Bob"]},
|
||||
)
|
||||
|
||||
# Replay from the checkpoint before sub_node in the 2nd invocation
|
||||
history = list(graph.get_state_history(config))
|
||||
before_sub = [s for s in history if s.next == ("sub_node",)]
|
||||
# First match is from the 2nd invocation (history is newest-first)
|
||||
before_sub_2nd = before_sub[0]
|
||||
|
||||
observed_sub_input.clear()
|
||||
replay = graph.invoke(None, before_sub_2nd.config)
|
||||
assert "__interrupt__" in replay
|
||||
|
||||
# Stateful subgraph retains state from 1st invocation, not from 2nd
|
||||
assert len(observed_sub_input) > 0
|
||||
step_a_state = observed_sub_input[0]
|
||||
assert step_a_state[0] == "step_a"
|
||||
assert step_a_state[1]["value"] == ["a:Alice", "b:30"]
|
||||
|
||||
|
||||
def test_stateful_subgraph_retains_state_on_parent_fork(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""With checkpointer=True, forking from the parent checkpoint before the
|
||||
2nd sub_node execution preserves the subgraph's accumulated state from
|
||||
the 1st invocation."""
|
||||
|
||||
observed_sub_input: list[tuple[str, dict]] = []
|
||||
|
||||
class SubState(TypedDict):
|
||||
value: Annotated[list[str], operator.add]
|
||||
|
||||
class ParentState(TypedDict):
|
||||
results: Annotated[list[str], operator.add]
|
||||
|
||||
def parent_node(state: ParentState) -> ParentState:
|
||||
return {"results": ["p"]}
|
||||
|
||||
def step_a(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_a", dict(state)))
|
||||
answer = interrupt("Question A")
|
||||
return {"value": [f"a:{answer}"]}
|
||||
|
||||
def step_b(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_b", dict(state)))
|
||||
answer = interrupt("Question B")
|
||||
return {"value": [f"b:{answer}"]}
|
||||
|
||||
sub = (
|
||||
StateGraph(SubState)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("step_b", step_b)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "step_b")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("parent_node", parent_node)
|
||||
.add_node("sub_node", sub)
|
||||
.add_edge(START, "parent_node")
|
||||
.add_edge("parent_node", "sub_node")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# 1st invocation: complete with 2 interrupts
|
||||
graph.invoke({"results": []}, config) # step_a interrupt
|
||||
graph.invoke(Command(resume="Alice"), config) # step_b interrupt
|
||||
graph.invoke(Command(resume="30"), config) # complete
|
||||
|
||||
# 2nd invocation: complete with 2 interrupts
|
||||
graph.invoke({"results": []}, config) # step_a interrupt
|
||||
graph.invoke(Command(resume="Bob"), config) # step_b interrupt
|
||||
graph.invoke(Command(resume="25"), config) # complete
|
||||
|
||||
# Fork from the checkpoint before sub_node in the 2nd invocation
|
||||
history = list(graph.get_state_history(config))
|
||||
before_sub = [s for s in history if s.next == ("sub_node",)]
|
||||
before_sub_2nd = before_sub[0]
|
||||
|
||||
fork_config = graph.update_state(before_sub_2nd.config, {"results": ["forked"]})
|
||||
|
||||
observed_sub_input.clear()
|
||||
fork_result = graph.invoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
|
||||
# Forked subgraph retains state from 1st invocation, not from 2nd
|
||||
assert len(observed_sub_input) > 0
|
||||
step_a_state = observed_sub_input[0]
|
||||
assert step_a_state[0] == "step_a"
|
||||
assert step_a_state[1]["value"] == ["a:Alice", "b:30"]
|
||||
|
||||
|
||||
def test_stateless_subgraph_starts_fresh_on_parent_replay(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Without checkpointer=True, the subgraph starts fresh when the parent
|
||||
replays. Graph: parent_node -> sub_node.
|
||||
|
||||
Invoke the graph twice on the same thread. Each invocation triggers
|
||||
two interrupts (step_a, step_b). After both complete, replay from the
|
||||
checkpoint before sub_node in the 2nd invocation. The stateless subgraph
|
||||
should see empty state because it has no persistent checkpoint history."""
|
||||
|
||||
observed_sub_input: list[tuple[str, dict]] = []
|
||||
|
||||
class SubState(TypedDict):
|
||||
value: Annotated[list[str], operator.add]
|
||||
|
||||
class ParentState(TypedDict):
|
||||
results: Annotated[list[str], operator.add]
|
||||
|
||||
def parent_node(state: ParentState) -> ParentState:
|
||||
return {"results": ["p"]}
|
||||
|
||||
def step_a(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_a", dict(state)))
|
||||
answer = interrupt("Question A")
|
||||
return {"value": [f"a:{answer}"]}
|
||||
|
||||
def step_b(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_b", dict(state)))
|
||||
answer = interrupt("Question B")
|
||||
return {"value": [f"b:{answer}"]}
|
||||
|
||||
sub = (
|
||||
StateGraph(SubState)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("step_b", step_b)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "step_b")
|
||||
.compile() # No checkpointer
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("parent_node", parent_node)
|
||||
.add_node("sub_node", sub)
|
||||
.add_edge(START, "parent_node")
|
||||
.add_edge("parent_node", "sub_node")
|
||||
.compile(checkpointer=sync_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# 1st invocation: complete with 2 interrupts
|
||||
graph.invoke({"results": []}, config) # step_a interrupt
|
||||
graph.invoke(Command(resume="Alice"), config) # step_b interrupt
|
||||
graph.invoke(Command(resume="30"), config) # complete
|
||||
|
||||
# Verify 1st invocation: subgraph started fresh
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": []})
|
||||
assert step_b_entries[0] == ("step_b", {"value": ["a:Alice"]})
|
||||
|
||||
# 2nd invocation: complete with 2 interrupts
|
||||
observed_sub_input.clear()
|
||||
graph.invoke({"results": []}, config) # step_a interrupt
|
||||
graph.invoke(Command(resume="Bob"), config) # step_b interrupt
|
||||
graph.invoke(Command(resume="25"), config) # complete
|
||||
|
||||
# Verify 2nd invocation: stateless subgraph starts fresh again
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": []})
|
||||
assert step_b_entries[0] == ("step_b", {"value": ["a:Bob"]})
|
||||
|
||||
# Replay from the checkpoint before sub_node in the 2nd invocation
|
||||
history = list(graph.get_state_history(config))
|
||||
before_sub = [s for s in history if s.next == ("sub_node",)]
|
||||
# First match is from the 2nd invocation (history is newest-first)
|
||||
before_sub_2nd = before_sub[0]
|
||||
|
||||
observed_sub_input.clear()
|
||||
replay = graph.invoke(None, before_sub_2nd.config)
|
||||
assert "__interrupt__" in replay
|
||||
|
||||
# Stateless subgraph starts fresh — no prior state
|
||||
assert len(observed_sub_input) > 0
|
||||
step_a_state = observed_sub_input[0]
|
||||
assert step_a_state[0] == "step_a"
|
||||
assert step_a_state[1]["value"] == [], (
|
||||
"Stateless subgraph should start fresh on replay"
|
||||
)
|
||||
|
||||
@@ -1587,3 +1587,288 @@ async def test_checkpoint_ns_accessible_in_subgraph(
|
||||
assert captured_config["checkpoint_ns"] is not None
|
||||
assert captured_config["checkpoint_ns"] != ""
|
||||
assert captured_config["thread_id"] == "1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Section 8: Stateful vs stateless subgraph state retention on replay
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ required for async test",
|
||||
)
|
||||
async def test_stateful_subgraph_retains_state_on_parent_replay(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""With checkpointer=True, the subgraph retains its accumulated state
|
||||
across parent invocations. Graph: parent_node -> sub_node.
|
||||
|
||||
Invoke the graph twice on the same thread. Each invocation triggers
|
||||
two interrupts (step_a, step_b). After both complete, replay from the
|
||||
checkpoint before sub_node in the 2nd invocation. The stateful subgraph
|
||||
should see accumulated state from the 1st invocation (a:Alice, b:30)
|
||||
but not the 2nd invocation's state (a:Bob, b:25)."""
|
||||
|
||||
observed_sub_input: list[tuple[str, dict]] = []
|
||||
|
||||
class SubState(TypedDict):
|
||||
value: Annotated[list[str], operator.add]
|
||||
|
||||
class ParentState(TypedDict):
|
||||
results: Annotated[list[str], operator.add]
|
||||
|
||||
def parent_node(state: ParentState) -> ParentState:
|
||||
return {"results": ["p"]}
|
||||
|
||||
def step_a(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_a", dict(state)))
|
||||
answer = interrupt("Question A")
|
||||
return {"value": [f"a:{answer}"]}
|
||||
|
||||
def step_b(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_b", dict(state)))
|
||||
answer = interrupt("Question B")
|
||||
return {"value": [f"b:{answer}"]}
|
||||
|
||||
sub = (
|
||||
StateGraph(SubState)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("step_b", step_b)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "step_b")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("parent_node", parent_node)
|
||||
.add_node("sub_node", sub)
|
||||
.add_edge(START, "parent_node")
|
||||
.add_edge("parent_node", "sub_node")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# 1st invocation: complete with 2 interrupts
|
||||
await graph.ainvoke({"results": []}, config) # step_a interrupt
|
||||
await graph.ainvoke(Command(resume="Alice"), config) # step_b interrupt
|
||||
await graph.ainvoke(Command(resume="30"), config) # complete
|
||||
|
||||
# Verify 1st invocation: subgraph started fresh
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": []})
|
||||
assert step_b_entries[0] == ("step_b", {"value": ["a:Alice"]})
|
||||
|
||||
# 2nd invocation: complete with 2 interrupts
|
||||
observed_sub_input.clear()
|
||||
await graph.ainvoke({"results": []}, config) # step_a interrupt
|
||||
await graph.ainvoke(Command(resume="Bob"), config) # step_b interrupt
|
||||
await graph.ainvoke(Command(resume="25"), config) # complete
|
||||
|
||||
# Verify 2nd invocation: subgraph retained state from 1st invocation
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": ["a:Alice", "b:30"]})
|
||||
assert step_b_entries[0] == (
|
||||
"step_b",
|
||||
{"value": ["a:Alice", "b:30", "a:Bob"]},
|
||||
)
|
||||
|
||||
# Replay from the checkpoint before sub_node in the 2nd invocation
|
||||
history = [s async for s in graph.aget_state_history(config)]
|
||||
before_sub = [s for s in history if s.next == ("sub_node",)]
|
||||
# First match is from the 2nd invocation (history is newest-first)
|
||||
before_sub_2nd = before_sub[0]
|
||||
|
||||
observed_sub_input.clear()
|
||||
replay = await graph.ainvoke(None, before_sub_2nd.config)
|
||||
assert "__interrupt__" in replay
|
||||
|
||||
# Stateful subgraph retains state from 1st invocation, not from 2nd
|
||||
assert len(observed_sub_input) > 0
|
||||
step_a_state = observed_sub_input[0]
|
||||
assert step_a_state[0] == "step_a"
|
||||
assert step_a_state[1]["value"] == ["a:Alice", "b:30"]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ required for async test",
|
||||
)
|
||||
async def test_stateful_subgraph_retains_state_on_parent_fork(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""With checkpointer=True, forking from the parent checkpoint before the
|
||||
2nd sub_node execution preserves the subgraph's accumulated state from
|
||||
the 1st invocation."""
|
||||
|
||||
observed_sub_input: list[tuple[str, dict]] = []
|
||||
|
||||
class SubState(TypedDict):
|
||||
value: Annotated[list[str], operator.add]
|
||||
|
||||
class ParentState(TypedDict):
|
||||
results: Annotated[list[str], operator.add]
|
||||
|
||||
def parent_node(state: ParentState) -> ParentState:
|
||||
return {"results": ["p"]}
|
||||
|
||||
def step_a(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_a", dict(state)))
|
||||
answer = interrupt("Question A")
|
||||
return {"value": [f"a:{answer}"]}
|
||||
|
||||
def step_b(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_b", dict(state)))
|
||||
answer = interrupt("Question B")
|
||||
return {"value": [f"b:{answer}"]}
|
||||
|
||||
sub = (
|
||||
StateGraph(SubState)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("step_b", step_b)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "step_b")
|
||||
.compile(checkpointer=True)
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("parent_node", parent_node)
|
||||
.add_node("sub_node", sub)
|
||||
.add_edge(START, "parent_node")
|
||||
.add_edge("parent_node", "sub_node")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# 1st invocation: complete with 2 interrupts
|
||||
await graph.ainvoke({"results": []}, config) # step_a interrupt
|
||||
await graph.ainvoke(Command(resume="Alice"), config) # step_b interrupt
|
||||
await graph.ainvoke(Command(resume="30"), config) # complete
|
||||
|
||||
# 2nd invocation: complete with 2 interrupts
|
||||
await graph.ainvoke({"results": []}, config) # step_a interrupt
|
||||
await graph.ainvoke(Command(resume="Bob"), config) # step_b interrupt
|
||||
await graph.ainvoke(Command(resume="25"), config) # complete
|
||||
|
||||
# Fork from the checkpoint before sub_node in the 2nd invocation
|
||||
history = [s async for s in graph.aget_state_history(config)]
|
||||
before_sub = [s for s in history if s.next == ("sub_node",)]
|
||||
before_sub_2nd = before_sub[0]
|
||||
|
||||
fork_config = await graph.aupdate_state(
|
||||
before_sub_2nd.config, {"results": ["forked"]}
|
||||
)
|
||||
|
||||
observed_sub_input.clear()
|
||||
fork_result = await graph.ainvoke(None, fork_config)
|
||||
assert "__interrupt__" in fork_result
|
||||
|
||||
# Forked subgraph retains state from 1st invocation, not from 2nd
|
||||
assert len(observed_sub_input) > 0
|
||||
step_a_state = observed_sub_input[0]
|
||||
assert step_a_state[0] == "step_a"
|
||||
assert step_a_state[1]["value"] == ["a:Alice", "b:30"]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ required for async test",
|
||||
)
|
||||
async def test_stateless_subgraph_starts_fresh_on_parent_replay(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
"""Without checkpointer=True, the subgraph starts fresh when the parent
|
||||
replays. Graph: parent_node -> sub_node.
|
||||
|
||||
Invoke the graph twice on the same thread. Each invocation triggers
|
||||
two interrupts (step_a, step_b). After both complete, replay from the
|
||||
checkpoint before sub_node in the 2nd invocation. The stateless subgraph
|
||||
should see empty state because it has no persistent checkpoint history."""
|
||||
|
||||
observed_sub_input: list[tuple[str, dict]] = []
|
||||
|
||||
class SubState(TypedDict):
|
||||
value: Annotated[list[str], operator.add]
|
||||
|
||||
class ParentState(TypedDict):
|
||||
results: Annotated[list[str], operator.add]
|
||||
|
||||
def parent_node(state: ParentState) -> ParentState:
|
||||
return {"results": ["p"]}
|
||||
|
||||
def step_a(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_a", dict(state)))
|
||||
answer = interrupt("Question A")
|
||||
return {"value": [f"a:{answer}"]}
|
||||
|
||||
def step_b(state: SubState) -> SubState:
|
||||
observed_sub_input.append(("step_b", dict(state)))
|
||||
answer = interrupt("Question B")
|
||||
return {"value": [f"b:{answer}"]}
|
||||
|
||||
sub = (
|
||||
StateGraph(SubState)
|
||||
.add_node("step_a", step_a)
|
||||
.add_node("step_b", step_b)
|
||||
.add_edge(START, "step_a")
|
||||
.add_edge("step_a", "step_b")
|
||||
.compile() # No checkpointer
|
||||
)
|
||||
|
||||
graph = (
|
||||
StateGraph(ParentState)
|
||||
.add_node("parent_node", parent_node)
|
||||
.add_node("sub_node", sub)
|
||||
.add_edge(START, "parent_node")
|
||||
.add_edge("parent_node", "sub_node")
|
||||
.compile(checkpointer=async_checkpointer)
|
||||
)
|
||||
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# 1st invocation: complete with 2 interrupts
|
||||
await graph.ainvoke({"results": []}, config) # step_a interrupt
|
||||
await graph.ainvoke(Command(resume="Alice"), config) # step_b interrupt
|
||||
await graph.ainvoke(Command(resume="30"), config) # complete
|
||||
|
||||
# Verify 1st invocation: subgraph started fresh
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": []})
|
||||
assert step_b_entries[0] == ("step_b", {"value": ["a:Alice"]})
|
||||
|
||||
# 2nd invocation: complete with 2 interrupts
|
||||
observed_sub_input.clear()
|
||||
await graph.ainvoke({"results": []}, config) # step_a interrupt
|
||||
await graph.ainvoke(Command(resume="Bob"), config) # step_b interrupt
|
||||
await graph.ainvoke(Command(resume="25"), config) # complete
|
||||
|
||||
# Verify 2nd invocation: stateless subgraph starts fresh again
|
||||
step_a_entries = [e for e in observed_sub_input if e[0] == "step_a"]
|
||||
step_b_entries = [e for e in observed_sub_input if e[0] == "step_b"]
|
||||
assert step_a_entries[0] == ("step_a", {"value": []})
|
||||
assert step_b_entries[0] == ("step_b", {"value": ["a:Bob"]})
|
||||
|
||||
# Replay from the checkpoint before sub_node in the 2nd invocation
|
||||
history = [s async for s in graph.aget_state_history(config)]
|
||||
before_sub = [s for s in history if s.next == ("sub_node",)]
|
||||
# First match is from the 2nd invocation (history is newest-first)
|
||||
before_sub_2nd = before_sub[0]
|
||||
|
||||
observed_sub_input.clear()
|
||||
replay = await graph.ainvoke(None, before_sub_2nd.config)
|
||||
assert "__interrupt__" in replay
|
||||
|
||||
# Stateless subgraph starts fresh — no prior state
|
||||
assert len(observed_sub_input) > 0
|
||||
step_a_state = observed_sub_input[0]
|
||||
assert step_a_state[0] == "step_a"
|
||||
assert step_a_state[1]["value"] == [], (
|
||||
"Stateless subgraph should start fresh on replay"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user