this is diabolical

This commit is contained in:
Sydney Runkle
2026-03-05 16:16:15 -08:00
parent 8ea279b8ad
commit 4826448e1a
3 changed files with 428 additions and 123 deletions
+103 -4
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
@@ -1106,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:
@@ -1115,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, []
@@ -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, []
+161 -61
View File
@@ -1568,12 +1568,14 @@ def test_checkpoint_ns_accessible_in_subgraph(
def test_stateful_subgraph_retains_state_on_parent_replay(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""With checkpointer=True, the subgraph retains its prior state when the
parent replays. Graph: parent_1 -> sub_1 -> parent_2 -> sub_2.
"""With checkpointer=True, the subgraph retains its accumulated state
across parent invocations. Graph: parent_node -> sub_node.
After completing the full flow, replay from before parent_2. The stateful
subgraph (sub_2) sees its old state from the prior execution because its
own checkpointer persists state across parent invocations."""
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]] = []
@@ -1583,11 +1585,8 @@ def test_stateful_subgraph_retains_state_on_parent_replay(
class ParentState(TypedDict):
results: Annotated[list[str], operator.add]
def parent_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
def parent_2(state: ParentState) -> ParentState:
return {"results": ["p2"]}
def parent_node(state: ParentState) -> ParentState:
return {"results": ["p"]}
def step_a(state: SubState) -> SubState:
observed_sub_input.append(("step_a", dict(state)))
@@ -1610,54 +1609,64 @@ def test_stateful_subgraph_retains_state_on_parent_replay(
graph = (
StateGraph(ParentState)
.add_node("parent_1", parent_1)
.add_node("sub_1", sub)
.add_node("parent_2", parent_2)
.add_node("sub_2", sub)
.add_edge(START, "parent_1")
.add_edge("parent_1", "sub_1")
.add_edge("sub_1", "parent_2")
.add_edge("parent_2", "sub_2")
.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"}}
# Complete the full flow (4 interrupts)
graph.invoke({"results": []}, config) # sub_1 step_a
graph.invoke(Command(resume="Alice"), config) # sub_1 step_b
graph.invoke(Command(resume="30"), config) # sub_2 step_a
graph.invoke(Command(resume="Bob"), config) # sub_2 step_b
graph.invoke(Command(resume="25"), config) # complete
# 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
# Replay from before parent_2
# 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_p2 = next(s for s in history if s.next == ("parent_2",))
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_p2.config)
replay = graph.invoke(None, before_sub_2nd.config)
assert "__interrupt__" in replay
# Stateful subgraph retains old state from prior sub_2 execution
# 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"] != [], (
"Stateful subgraph should retain prior state on replay"
)
assert "a:Bob" in step_a_state[1]["value"]
assert "b:25" in step_a_state[1]["value"]
assert step_a_state[1]["value"] == ["a:Alice", "b:30"]
def test_stateless_subgraph_starts_fresh_on_parent_replay(
def test_stateful_subgraph_retains_state_on_parent_fork(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
"""Without checkpointer=True, the subgraph starts fresh when the parent
replays. Graph: parent_1 -> sub_1 -> parent_2 -> sub_2.
After completing the full flow, replay from before parent_2. The stateless
subgraph (sub_2) sees empty state because it has no persistent checkpoint
history of its own."""
"""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]] = []
@@ -1667,11 +1676,88 @@ def test_stateless_subgraph_starts_fresh_on_parent_replay(
class ParentState(TypedDict):
results: Annotated[list[str], operator.add]
def parent_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
def parent_node(state: ParentState) -> ParentState:
return {"results": ["p"]}
def parent_2(state: ParentState) -> ParentState:
return {"results": ["p2"]}
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)))
@@ -1694,32 +1780,46 @@ def test_stateless_subgraph_starts_fresh_on_parent_replay(
graph = (
StateGraph(ParentState)
.add_node("parent_1", parent_1)
.add_node("sub_1", sub)
.add_node("parent_2", parent_2)
.add_node("sub_2", sub)
.add_edge(START, "parent_1")
.add_edge("parent_1", "sub_1")
.add_edge("sub_1", "parent_2")
.add_edge("parent_2", "sub_2")
.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"}}
# Complete the full flow (4 interrupts)
graph.invoke({"results": []}, config) # sub_1 step_a
graph.invoke(Command(resume="Alice"), config) # sub_1 step_b
graph.invoke(Command(resume="30"), config) # sub_2 step_a
graph.invoke(Command(resume="Bob"), config) # sub_2 step_b
graph.invoke(Command(resume="25"), config) # complete
# 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
# Replay from before parent_2
# 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_p2 = next(s for s in history if s.next == ("parent_2",))
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_p2.config)
replay = graph.invoke(None, before_sub_2nd.config)
assert "__interrupt__" in replay
# Stateless subgraph starts fresh — no prior state
+164 -58
View File
@@ -1601,12 +1601,14 @@ async def test_checkpoint_ns_accessible_in_subgraph(
async def test_stateful_subgraph_retains_state_on_parent_replay(
async_checkpointer: BaseCheckpointSaver,
) -> None:
"""With checkpointer=True, the subgraph retains its prior state when the
parent replays. Graph: parent_1 -> sub_1 -> parent_2 -> sub_2.
"""With checkpointer=True, the subgraph retains its accumulated state
across parent invocations. Graph: parent_node -> sub_node.
After completing the full flow, replay from before parent_2. The stateful
subgraph (sub_2) sees its old state from the prior execution because its
own checkpointer persists state across parent invocations."""
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]] = []
@@ -1616,11 +1618,8 @@ async def test_stateful_subgraph_retains_state_on_parent_replay(
class ParentState(TypedDict):
results: Annotated[list[str], operator.add]
def parent_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
def parent_2(state: ParentState) -> ParentState:
return {"results": ["p2"]}
def parent_node(state: ParentState) -> ParentState:
return {"results": ["p"]}
def step_a(state: SubState) -> SubState:
observed_sub_input.append(("step_a", dict(state)))
@@ -1643,43 +1642,138 @@ async def test_stateful_subgraph_retains_state_on_parent_replay(
graph = (
StateGraph(ParentState)
.add_node("parent_1", parent_1)
.add_node("sub_1", sub)
.add_node("parent_2", parent_2)
.add_node("sub_2", sub)
.add_edge(START, "parent_1")
.add_edge("parent_1", "sub_1")
.add_edge("sub_1", "parent_2")
.add_edge("parent_2", "sub_2")
.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"}}
# Complete the full flow (4 interrupts)
await graph.ainvoke({"results": []}, config) # sub_1 step_a
await graph.ainvoke(Command(resume="Alice"), config) # sub_1 step_b
await graph.ainvoke(Command(resume="30"), config) # sub_2 step_a
await graph.ainvoke(Command(resume="Bob"), config) # sub_2 step_b
await graph.ainvoke(Command(resume="25"), config) # complete
# 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
# Replay from before parent_2
# 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_p2 = next(s for s in history if s.next == ("parent_2",))
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_p2.config)
replay = await graph.ainvoke(None, before_sub_2nd.config)
assert "__interrupt__" in replay
# Stateful subgraph retains old state from prior sub_2 execution
# 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"] != [], (
"Stateful subgraph should retain prior state on replay"
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)
)
assert "a:Bob" in step_a_state[1]["value"]
assert "b:25" in step_a_state[1]["value"]
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(
@@ -1690,11 +1784,12 @@ 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_1 -> sub_1 -> parent_2 -> sub_2.
replays. Graph: parent_node -> sub_node.
After completing the full flow, replay from before parent_2. The stateless
subgraph (sub_2) sees empty state because it has no persistent checkpoint
history of its own."""
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]] = []
@@ -1704,11 +1799,8 @@ async def test_stateless_subgraph_starts_fresh_on_parent_replay(
class ParentState(TypedDict):
results: Annotated[list[str], operator.add]
def parent_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
def parent_2(state: ParentState) -> ParentState:
return {"results": ["p2"]}
def parent_node(state: ParentState) -> ParentState:
return {"results": ["p"]}
def step_a(state: SubState) -> SubState:
observed_sub_input.append(("step_a", dict(state)))
@@ -1731,32 +1823,46 @@ async def test_stateless_subgraph_starts_fresh_on_parent_replay(
graph = (
StateGraph(ParentState)
.add_node("parent_1", parent_1)
.add_node("sub_1", sub)
.add_node("parent_2", parent_2)
.add_node("sub_2", sub)
.add_edge(START, "parent_1")
.add_edge("parent_1", "sub_1")
.add_edge("sub_1", "parent_2")
.add_edge("parent_2", "sub_2")
.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"}}
# Complete the full flow (4 interrupts)
await graph.ainvoke({"results": []}, config) # sub_1 step_a
await graph.ainvoke(Command(resume="Alice"), config) # sub_1 step_b
await graph.ainvoke(Command(resume="30"), config) # sub_2 step_a
await graph.ainvoke(Command(resume="Bob"), config) # sub_2 step_b
await graph.ainvoke(Command(resume="25"), config) # complete
# 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
# Replay from before parent_2
# 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_p2 = next(s for s in history if s.next == ("parent_2",))
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_p2.config)
replay = await graph.ainvoke(None, before_sub_2nd.config)
assert "__interrupt__" in replay
# Stateless subgraph starts fresh — no prior state