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
3 changed files with 673 additions and 32 deletions
+117 -32
View File
@@ -245,9 +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.is_replaying = CONFIG_KEY_CHECKPOINT_ID in config[
self.is_replaying = CONFIG_KEY_CHECKPOINT_ID in config[CONF] or config[
CONF
] or config[CONF].get(CONFIG_KEY_REPLAYING, False)
].get(CONFIG_KEY_REPLAYING, False)
self._migrate_checkpoint = migrate_checkpoint
self.trigger_to_nodes = trigger_to_nodes
self.retry_policy = retry_policy
@@ -645,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:
@@ -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:
@@ -1101,12 +1168,14 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
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 = None
saved = self._get_checkpoint_before_parent()
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1131,20 +1200,6 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
if saved.pending_writes is not None
else []
)
# 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=...) on a specific checkpoint), keep them —
# multi-interrupt scenarios need previously resolved values preserved.
if self.is_replaying:
has_resume_value = (
isinstance(self.input, Command) and self.input.resume is not None
) or self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING, False)
if 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
@@ -1292,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:
@@ -1301,12 +1398,14 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
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 = None
saved = await self._aget_checkpoint_before_parent()
if saved is None:
saved = CheckpointTuple(
self.checkpoint_config, empty_checkpoint(), {"step": -2}, None, []
@@ -1331,20 +1430,6 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
if saved.pending_writes is not None
else []
)
# 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=...) on a specific checkpoint), keep them —
# multi-interrupt scenarios need previously resolved values preserved.
if self.is_replaying:
has_resume_value = (
isinstance(self.input, Command) and self.input.resume is not None
) or self.config.get(CONF, {}).get(CONFIG_KEY_RESUMING, False)
if 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)
)
+271
View File
@@ -1558,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"
)