mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-10 03:37:51 +02:00
this is diabolical
This commit is contained in:
@@ -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, []
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user