This commit is contained in:
Sydney Runkle
2026-03-05 11:04:13 -08:00
parent eb2f09e321
commit 8ea279b8ad
2 changed files with 350 additions and 0 deletions
+171
View File
@@ -1558,3 +1558,174 @@ 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 prior state when the
parent replays. Graph: parent_1 -> sub_1 -> parent_2 -> sub_2.
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."""
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_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
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_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")
.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
# Replay from before parent_2
history = list(graph.get_state_history(config))
before_p2 = next(s for s in history if s.next == ("parent_2",))
observed_sub_input.clear()
replay = graph.invoke(None, before_p2.config)
assert "__interrupt__" in replay
# Stateful subgraph retains old state from prior sub_2 execution
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"]
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_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."""
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_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
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() # No checkpointer
)
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")
.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
# Replay from before parent_2
history = list(graph.get_state_history(config))
before_p2 = next(s for s in history if s.next == ("parent_2",))
observed_sub_input.clear()
replay = graph.invoke(None, before_p2.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,182 @@ 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 prior state when the
parent replays. Graph: parent_1 -> sub_1 -> parent_2 -> sub_2.
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."""
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_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
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_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")
.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
# Replay from before parent_2
history = [s async for s in graph.aget_state_history(config)]
before_p2 = next(s for s in history if s.next == ("parent_2",))
observed_sub_input.clear()
replay = await graph.ainvoke(None, before_p2.config)
assert "__interrupt__" in replay
# Stateful subgraph retains old state from prior sub_2 execution
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"]
@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_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."""
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_1(state: ParentState) -> ParentState:
return {"results": ["p1"]}
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() # No checkpointer
)
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")
.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
# Replay from before parent_2
history = [s async for s in graph.aget_state_history(config)]
before_p2 = next(s for s in history if s.next == ("parent_2",))
observed_sub_input.clear()
replay = await graph.ainvoke(None, before_p2.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"
)