Compare commits

...
Author SHA1 Message Date
Sydney Runkle 4826448e1a this is diabolical 2026-03-05 16:16:15 -08:00
Sydney Runkle 8ea279b8ad tests 2026-03-05 11:04:13 -08:00
Sydney Runkle eb2f09e321 push 2026-03-05 09:51:24 -08:00
Sydney Runkle 26d279a0ac rename 2026-03-05 09:30:02 -08:00
Sydney Runkle e850b21d08 continue 2026-03-05 09:10:43 -08:00
Sydney Runkle 6a92b7ff3c alt fix idea 2026-03-05 08:54:10 -08:00
4 changed files with 716 additions and 46 deletions
@@ -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,
+146 -45
View File
@@ -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)
)
+282 -1
View File
@@ -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"
)