mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7d601b1b82 | ||
|
|
314cb19da8 | ||
|
|
0207ec7cff | ||
|
|
a4a028a254 | ||
|
|
0afdb91adc |
@@ -226,6 +226,20 @@ def _needs_replay(spec: BaseChannel, stored: object) -> bool:
|
||||
return stored is MISSING
|
||||
|
||||
|
||||
def _require_saver_for_history(
|
||||
checkpoint: Checkpoint,
|
||||
delta_channels: list[str],
|
||||
saver: BaseCheckpointSaver | None,
|
||||
config: RunnableConfig | None,
|
||||
) -> None:
|
||||
written = [k for k in delta_channels if k in checkpoint["channel_versions"]]
|
||||
if written and (saver is None or config is None):
|
||||
raise ValueError(
|
||||
f"DeltaChannel {written} has history to replay but no checkpointer "
|
||||
"or config was passed to read it"
|
||||
)
|
||||
|
||||
|
||||
def channels_from_checkpoint(
|
||||
specs: Mapping[str, BaseChannel | ManagedValueSpec],
|
||||
checkpoint: Checkpoint,
|
||||
@@ -256,6 +270,7 @@ def channels_from_checkpoint(
|
||||
for k, spec in channel_specs.items()
|
||||
if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING))
|
||||
]
|
||||
_require_saver_for_history(checkpoint, delta_channels, saver, config)
|
||||
histories: Mapping[str, Any] = {}
|
||||
if delta_channels and saver is not None and config is not None:
|
||||
histories = saver.get_delta_channel_history(
|
||||
@@ -298,6 +313,7 @@ async def achannels_from_checkpoint(
|
||||
for k, spec in channel_specs.items()
|
||||
if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING))
|
||||
]
|
||||
_require_saver_for_history(checkpoint, delta_channels, saver, config)
|
||||
histories: Mapping[str, Any] = {}
|
||||
if delta_channels and saver is not None and config is not None:
|
||||
histories = await saver.aget_delta_channel_history(
|
||||
|
||||
@@ -835,6 +835,41 @@ class Pregel(
|
||||
if auto_validate:
|
||||
self.validate()
|
||||
|
||||
def _resolve_checkpointer(
|
||||
self, config: RunnableConfig
|
||||
) -> BaseCheckpointSaver | None:
|
||||
"""The saver runs and state methods use: none for `checkpointer=False`,
|
||||
else the one a parent lends a subgraph through the config, else this
|
||||
graph's own."""
|
||||
if self.checkpointer is False:
|
||||
return None
|
||||
conf = config.get(CONF, {})
|
||||
if CONFIG_KEY_CHECKPOINTER in conf:
|
||||
checkpointer = conf[CONFIG_KEY_CHECKPOINTER]
|
||||
elif self.checkpointer is True:
|
||||
raise RuntimeError("checkpointer=True cannot be used for root graphs.")
|
||||
else:
|
||||
checkpointer = self.checkpointer
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
return checkpointer
|
||||
|
||||
def _state_checkpointer(self, config: RunnableConfig) -> BaseCheckpointSaver:
|
||||
checkpointer = self._resolve_checkpointer(ensure_config(config))
|
||||
if not isinstance(checkpointer, BaseCheckpointSaver):
|
||||
raise ValueError("No checkpointer set")
|
||||
return checkpointer
|
||||
|
||||
def _own_checkpoint_config(self, config: RunnableConfig) -> RunnableConfig:
|
||||
"""A `checkpointer=True` subgraph keeps one history per thread, stored
|
||||
under its namespace with the task ids removed."""
|
||||
if self.checkpointer is not True:
|
||||
return config
|
||||
ns = config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
return patch_configurable(
|
||||
config, {CONFIG_KEY_CHECKPOINT_NS: recast_checkpoint_ns(ns)}
|
||||
)
|
||||
|
||||
def _apply_checkpointer_allowlist(
|
||||
self, checkpointer: BaseCheckpointSaver | None
|
||||
) -> BaseCheckpointSaver | None:
|
||||
@@ -1146,7 +1181,9 @@ class Pregel(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
recurse: BaseCheckpointSaver | None = None,
|
||||
*,
|
||||
saver: BaseCheckpointSaver,
|
||||
recurse: bool = False,
|
||||
apply_pending_writes: bool = False,
|
||||
) -> StateSnapshot:
|
||||
if not saved:
|
||||
@@ -1169,9 +1206,7 @@ class Pregel(
|
||||
channels, managed = channels_from_checkpoint(
|
||||
self.channels,
|
||||
saved.checkpoint,
|
||||
saver=self.checkpointer
|
||||
if isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else None,
|
||||
saver=saver,
|
||||
config=saved.config,
|
||||
)
|
||||
# tasks for this checkpoint
|
||||
@@ -1186,11 +1221,7 @@ class Pregel(
|
||||
stop,
|
||||
for_execution=True,
|
||||
store=self.store,
|
||||
checkpointer=(
|
||||
self.checkpointer
|
||||
if isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else None
|
||||
),
|
||||
checkpointer=saver,
|
||||
manager=None,
|
||||
)
|
||||
# get the subgraphs
|
||||
@@ -1217,7 +1248,7 @@ class Pregel(
|
||||
# get the state of the subgraph
|
||||
config = {
|
||||
CONF: {
|
||||
CONFIG_KEY_CHECKPOINTER: recurse,
|
||||
CONFIG_KEY_CHECKPOINTER: saver,
|
||||
"thread_id": saved.config[CONF]["thread_id"],
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_ns,
|
||||
}
|
||||
@@ -1269,7 +1300,9 @@ class Pregel(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
saved: CheckpointTuple | None,
|
||||
recurse: BaseCheckpointSaver | None = None,
|
||||
*,
|
||||
saver: BaseCheckpointSaver,
|
||||
recurse: bool = False,
|
||||
apply_pending_writes: bool = False,
|
||||
) -> StateSnapshot:
|
||||
if not saved:
|
||||
@@ -1292,9 +1325,7 @@ class Pregel(
|
||||
channels, managed = await achannels_from_checkpoint(
|
||||
self.channels,
|
||||
saved.checkpoint,
|
||||
saver=self.checkpointer
|
||||
if isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else None,
|
||||
saver=saver,
|
||||
config=saved.config,
|
||||
)
|
||||
# tasks for this checkpoint
|
||||
@@ -1309,11 +1340,7 @@ class Pregel(
|
||||
stop,
|
||||
for_execution=True,
|
||||
store=self.store,
|
||||
checkpointer=(
|
||||
self.checkpointer
|
||||
if isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else None
|
||||
),
|
||||
checkpointer=saver,
|
||||
manager=None,
|
||||
)
|
||||
# get the subgraphs
|
||||
@@ -1340,7 +1367,7 @@ class Pregel(
|
||||
# get the state of the subgraph
|
||||
config = {
|
||||
CONF: {
|
||||
CONFIG_KEY_CHECKPOINTER: recurse,
|
||||
CONFIG_KEY_CHECKPOINTER: saver,
|
||||
"thread_id": saved.config[CONF]["thread_id"],
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_ns,
|
||||
}
|
||||
@@ -1393,13 +1420,7 @@ class Pregel(
|
||||
self, config: RunnableConfig, *, subgraphs: bool = False
|
||||
) -> StateSnapshot:
|
||||
"""Get the current state of the graph."""
|
||||
checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
if not checkpointer:
|
||||
raise ValueError("No checkpointer set")
|
||||
checkpointer = self._state_checkpointer(config)
|
||||
|
||||
if (
|
||||
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
@@ -1416,11 +1437,7 @@ class Pregel(
|
||||
raise ValueError(f"Subgraph {recast} not found")
|
||||
|
||||
config = merge_configs(self.config, config) if self.config else config
|
||||
if self.checkpointer is True:
|
||||
ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS])
|
||||
config = merge_configs(
|
||||
config, {CONF: {CONFIG_KEY_CHECKPOINT_NS: recast_checkpoint_ns(ns)}}
|
||||
)
|
||||
config = self._own_checkpoint_config(config)
|
||||
thread_id = config[CONF][CONFIG_KEY_THREAD_ID]
|
||||
if not isinstance(thread_id, str):
|
||||
config[CONF][CONFIG_KEY_THREAD_ID] = str(thread_id)
|
||||
@@ -1429,7 +1446,8 @@ class Pregel(
|
||||
return self._prepare_state_snapshot(
|
||||
config,
|
||||
saved,
|
||||
recurse=checkpointer if subgraphs else None,
|
||||
saver=checkpointer,
|
||||
recurse=subgraphs,
|
||||
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
)
|
||||
|
||||
@@ -1437,13 +1455,7 @@ class Pregel(
|
||||
self, config: RunnableConfig, *, subgraphs: bool = False
|
||||
) -> StateSnapshot:
|
||||
"""Get the current state of the graph."""
|
||||
checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
if not checkpointer:
|
||||
raise ValueError("No checkpointer set")
|
||||
checkpointer = self._state_checkpointer(config)
|
||||
|
||||
if (
|
||||
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
@@ -1460,11 +1472,7 @@ class Pregel(
|
||||
raise ValueError(f"Subgraph {recast} not found")
|
||||
|
||||
config = merge_configs(self.config, config) if self.config else config
|
||||
if self.checkpointer is True:
|
||||
ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS])
|
||||
config = merge_configs(
|
||||
config, {CONF: {CONFIG_KEY_CHECKPOINT_NS: recast_checkpoint_ns(ns)}}
|
||||
)
|
||||
config = self._own_checkpoint_config(config)
|
||||
thread_id = config[CONF][CONFIG_KEY_THREAD_ID]
|
||||
if not isinstance(thread_id, str):
|
||||
config[CONF][CONFIG_KEY_THREAD_ID] = str(thread_id)
|
||||
@@ -1473,7 +1481,8 @@ class Pregel(
|
||||
return await self._aprepare_state_snapshot(
|
||||
config,
|
||||
saved,
|
||||
recurse=checkpointer if subgraphs else None,
|
||||
saver=checkpointer,
|
||||
recurse=subgraphs,
|
||||
apply_pending_writes=CONFIG_KEY_CHECKPOINT_ID not in config[CONF],
|
||||
)
|
||||
|
||||
@@ -1487,13 +1496,7 @@ class Pregel(
|
||||
) -> Iterator[StateSnapshot]:
|
||||
"""Get the history of the state of the graph."""
|
||||
config = ensure_config(config)
|
||||
checkpointer: BaseCheckpointSaver | None = config[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
if not checkpointer:
|
||||
raise ValueError("No checkpointer set")
|
||||
checkpointer = self._state_checkpointer(config)
|
||||
|
||||
if (
|
||||
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
@@ -1522,12 +1525,13 @@ class Pregel(
|
||||
}
|
||||
},
|
||||
)
|
||||
config = self._own_checkpoint_config(config)
|
||||
# eagerly consume list() to avoid holding up the db cursor
|
||||
for checkpoint_tuple in list(
|
||||
checkpointer.list(config, before=before, limit=limit, filter=filter)
|
||||
):
|
||||
yield self._prepare_state_snapshot(
|
||||
checkpoint_tuple.config, checkpoint_tuple
|
||||
checkpoint_tuple.config, checkpoint_tuple, saver=checkpointer
|
||||
)
|
||||
|
||||
async def aget_state_history(
|
||||
@@ -1540,13 +1544,7 @@ class Pregel(
|
||||
) -> AsyncIterator[StateSnapshot]:
|
||||
"""Asynchronously get the history of the state of the graph."""
|
||||
config = ensure_config(config)
|
||||
checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
if not checkpointer:
|
||||
raise ValueError("No checkpointer set")
|
||||
checkpointer = self._state_checkpointer(config)
|
||||
|
||||
if (
|
||||
checkpoint_ns := config[CONF].get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
@@ -1576,6 +1574,7 @@ class Pregel(
|
||||
}
|
||||
},
|
||||
)
|
||||
config = self._own_checkpoint_config(config)
|
||||
# eagerly consume list() to avoid holding up the db cursor
|
||||
for checkpoint_tuple in [
|
||||
c
|
||||
@@ -1584,7 +1583,7 @@ class Pregel(
|
||||
)
|
||||
]:
|
||||
yield await self._aprepare_state_snapshot(
|
||||
checkpoint_tuple.config, checkpoint_tuple
|
||||
checkpoint_tuple.config, checkpoint_tuple, saver=checkpointer
|
||||
)
|
||||
|
||||
def bulk_update_state(
|
||||
@@ -1608,13 +1607,7 @@ class Pregel(
|
||||
RunnableConfig: The updated config.
|
||||
"""
|
||||
|
||||
checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
if not checkpointer:
|
||||
raise ValueError("No checkpointer set")
|
||||
checkpointer = self._state_checkpointer(config)
|
||||
|
||||
if len(supersteps) == 0:
|
||||
raise ValueError("No supersteps provided")
|
||||
@@ -1641,7 +1634,9 @@ class Pregel(
|
||||
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
||||
) -> RunnableConfig:
|
||||
# get last checkpoint
|
||||
config = ensure_config(self.config, input_config)
|
||||
config = self._own_checkpoint_config(
|
||||
ensure_config(self.config, input_config)
|
||||
)
|
||||
saved = checkpointer.get_tuple(config)
|
||||
if saved is not None:
|
||||
self._migrate_checkpoint(saved.checkpoint)
|
||||
@@ -1666,10 +1661,7 @@ class Pregel(
|
||||
channels, managed = channels_from_checkpoint(
|
||||
self.channels,
|
||||
checkpoint,
|
||||
saver=self.checkpointer
|
||||
if saved is not None
|
||||
and isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else None,
|
||||
saver=checkpointer,
|
||||
config=saved.config if saved is not None else None,
|
||||
)
|
||||
values, as_node = updates[0][:2]
|
||||
@@ -2074,13 +2066,7 @@ class Pregel(
|
||||
RunnableConfig: The updated config.
|
||||
"""
|
||||
|
||||
checkpointer: BaseCheckpointSaver | None = ensure_config(config)[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINTER, self.checkpointer
|
||||
)
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
if not checkpointer:
|
||||
raise ValueError("No checkpointer set")
|
||||
checkpointer = self._state_checkpointer(config)
|
||||
|
||||
if len(supersteps) == 0:
|
||||
raise ValueError("No supersteps provided")
|
||||
@@ -2107,7 +2093,9 @@ class Pregel(
|
||||
input_config: RunnableConfig, updates: Sequence[StateUpdate]
|
||||
) -> RunnableConfig:
|
||||
# get last checkpoint
|
||||
config = ensure_config(self.config, input_config)
|
||||
config = self._own_checkpoint_config(
|
||||
ensure_config(self.config, input_config)
|
||||
)
|
||||
saved = await checkpointer.aget_tuple(config)
|
||||
if saved is not None:
|
||||
self._migrate_checkpoint(saved.checkpoint)
|
||||
@@ -2132,10 +2120,7 @@ class Pregel(
|
||||
channels, managed = await achannels_from_checkpoint(
|
||||
self.channels,
|
||||
checkpoint,
|
||||
saver=self.checkpointer
|
||||
if saved is not None
|
||||
and isinstance(self.checkpointer, BaseCheckpointSaver)
|
||||
else None,
|
||||
saver=checkpointer,
|
||||
config=saved.config if saved is not None else None,
|
||||
)
|
||||
values, as_node = updates[0][:2]
|
||||
@@ -2576,16 +2561,7 @@ class Pregel(
|
||||
stream_modes.add(print_mode)
|
||||
else:
|
||||
stream_modes.update(print_mode)
|
||||
if self.checkpointer is False:
|
||||
checkpointer: BaseCheckpointSaver | None = None
|
||||
elif CONFIG_KEY_CHECKPOINTER in config.get(CONF, {}):
|
||||
checkpointer = config[CONF][CONFIG_KEY_CHECKPOINTER]
|
||||
elif self.checkpointer is True:
|
||||
raise RuntimeError("checkpointer=True cannot be used for root graphs.")
|
||||
else:
|
||||
checkpointer = self.checkpointer
|
||||
if isinstance(checkpointer, BaseCheckpointSaver):
|
||||
checkpointer = self._apply_checkpointer_allowlist(checkpointer)
|
||||
checkpointer = self._resolve_checkpointer(config)
|
||||
if checkpointer and not config.get(CONF):
|
||||
raise ValueError(
|
||||
"Checkpointer requires one or more of the following 'configurable' "
|
||||
@@ -2804,9 +2780,7 @@ class Pregel(
|
||||
"`durability` has no effect when no checkpointer is present.",
|
||||
)
|
||||
# set up subgraph checkpointing
|
||||
if self.checkpointer is True:
|
||||
ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS])
|
||||
config[CONF][CONFIG_KEY_CHECKPOINT_NS] = recast_checkpoint_ns(ns)
|
||||
config = self._own_checkpoint_config(config)
|
||||
# set up messages stream mode
|
||||
if "messages" in stream_modes:
|
||||
ns_ = cast(str | None, config[CONF].get(CONFIG_KEY_CHECKPOINT_NS))
|
||||
@@ -3231,9 +3205,7 @@ class Pregel(
|
||||
"`durability` has no effect when no checkpointer is present.",
|
||||
)
|
||||
# set up subgraph checkpointing
|
||||
if self.checkpointer is True:
|
||||
ns = cast(str, config[CONF][CONFIG_KEY_CHECKPOINT_NS])
|
||||
config[CONF][CONFIG_KEY_CHECKPOINT_NS] = recast_checkpoint_ns(ns)
|
||||
config = self._own_checkpoint_config(config)
|
||||
# set up messages stream mode
|
||||
if "messages" in stream_modes:
|
||||
# namespace can be None in a root level graph?
|
||||
|
||||
@@ -0,0 +1,395 @@
|
||||
import operator
|
||||
from typing import Annotated, Any, Literal
|
||||
|
||||
import pytest
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph._internal._constants import CONFIG_KEY_CHECKPOINTER
|
||||
from langgraph.channels.delta import DeltaChannel
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.pregel._checkpoint import (
|
||||
achannels_from_checkpoint,
|
||||
channels_from_checkpoint,
|
||||
empty_checkpoint,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
def _extend(state: list | None, writes: list[Any]) -> list:
|
||||
out = list(state or [])
|
||||
for write in writes:
|
||||
out.extend(write if isinstance(write, list) else [write])
|
||||
return out
|
||||
|
||||
|
||||
def _state_schema(snapshot_frequency: int = 1000) -> type:
|
||||
class State(TypedDict, total=False):
|
||||
delta: Annotated[
|
||||
list, DeltaChannel(_extend, snapshot_frequency=snapshot_frequency)
|
||||
]
|
||||
plain: Annotated[list, operator.add]
|
||||
|
||||
return State
|
||||
|
||||
|
||||
def _both(*items: str) -> dict:
|
||||
return {"delta": list(items), "plain": list(items)}
|
||||
|
||||
|
||||
def _child_builder(*, snapshot_frequency: int = 1000) -> StateGraph:
|
||||
builder = StateGraph(_state_schema(snapshot_frequency))
|
||||
builder.add_node("a", lambda state: _both("a1"))
|
||||
builder.add_node("b", lambda state: _both("b1", "b2"))
|
||||
builder.add_edge(START, "a")
|
||||
builder.add_edge("a", "b")
|
||||
builder.add_edge("b", END)
|
||||
return builder
|
||||
|
||||
|
||||
def _wrap(
|
||||
inner: StateGraph,
|
||||
*,
|
||||
checkpointer: bool | None = None,
|
||||
interrupt_before: list[str] | None = None,
|
||||
) -> StateGraph:
|
||||
builder = StateGraph(inner.state_schema)
|
||||
builder.add_node(
|
||||
"child",
|
||||
inner.compile(checkpointer=checkpointer, interrupt_before=interrupt_before),
|
||||
)
|
||||
builder.add_edge(START, "child")
|
||||
builder.add_edge("child", END)
|
||||
return builder
|
||||
|
||||
|
||||
def _nested_app(
|
||||
checkpointer: BaseCheckpointSaver,
|
||||
*,
|
||||
depth: int = 1,
|
||||
snapshot_frequency: int = 1000,
|
||||
pause_before_b: bool = False,
|
||||
subgraph_checkpointer: bool | None = None,
|
||||
) -> Any:
|
||||
graph = _child_builder(snapshot_frequency=snapshot_frequency)
|
||||
for _ in range(depth):
|
||||
graph = _wrap(
|
||||
graph,
|
||||
checkpointer=subgraph_checkpointer,
|
||||
interrupt_before=["b"] if pause_before_b else None,
|
||||
)
|
||||
return graph.compile(checkpointer=checkpointer)
|
||||
|
||||
|
||||
def _scoped(config: dict, namespace: str) -> dict:
|
||||
return {"configurable": {**config["configurable"], "checkpoint_ns": namespace}}
|
||||
|
||||
|
||||
def _child_namespace(app: Any, config: dict, *, depth: int = 1) -> str:
|
||||
namespace = ""
|
||||
for level in range(depth):
|
||||
scoped = _scoped(config, namespace) if namespace else config
|
||||
namespace = next(
|
||||
(
|
||||
task.state["configurable"]["checkpoint_ns"]
|
||||
for snapshot in app.get_state_history(scoped)
|
||||
for task in snapshot.tasks
|
||||
if task.name == "child" and isinstance(task.state, dict)
|
||||
),
|
||||
"",
|
||||
)
|
||||
assert namespace, f"no `child` subgraph task at nesting level {level}"
|
||||
return namespace
|
||||
|
||||
|
||||
async def _achild_namespace(app: Any, config: dict) -> str:
|
||||
async for snapshot in app.aget_state_history(config):
|
||||
for task in snapshot.tasks:
|
||||
if task.name == "child" and isinstance(task.state, dict):
|
||||
return task.state["configurable"]["checkpoint_ns"]
|
||||
raise AssertionError("no `child` subgraph task")
|
||||
|
||||
|
||||
HISTORY = [_both("a1", "b1", "b2"), _both("a1"), _both(), _both()]
|
||||
|
||||
|
||||
def test_subgraph_get_state(sync_checkpointer: BaseCheckpointSaver) -> None:
|
||||
app = _nested_app(sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
child = _scoped(config, _child_namespace(app, config))
|
||||
|
||||
assert app.get_state(config).values == _both("a1", "b1", "b2")
|
||||
assert app.get_state(child).values == _both("a1", "b1", "b2")
|
||||
|
||||
|
||||
async def test_subgraph_aget_state(async_checkpointer: BaseCheckpointSaver) -> None:
|
||||
app = _nested_app(async_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await app.ainvoke({}, config)
|
||||
|
||||
child = _scoped(config, await _achild_namespace(app, config))
|
||||
|
||||
assert (await app.aget_state(child)).values == _both("a1", "b1", "b2")
|
||||
|
||||
|
||||
def test_subgraph_get_state_history(sync_checkpointer: BaseCheckpointSaver) -> None:
|
||||
app = _nested_app(sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
child = _scoped(config, _child_namespace(app, config))
|
||||
|
||||
assert [s.values for s in app.get_state_history(child)] == HISTORY
|
||||
|
||||
|
||||
async def test_subgraph_aget_state_history(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
app = _nested_app(async_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await app.ainvoke({}, config)
|
||||
|
||||
child = _scoped(config, await _achild_namespace(app, config))
|
||||
|
||||
assert [s.values async for s in app.aget_state_history(child)] == HISTORY
|
||||
|
||||
|
||||
def test_doubly_nested_subgraph_get_state(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
app = _nested_app(sync_checkpointer, depth=2)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
child = _scoped(config, _child_namespace(app, config, depth=2))
|
||||
|
||||
assert app.get_state(child).values == _both("a1", "b1", "b2")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
|
||||
def test_interrupted_subgraph_task_state(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
persistence: Literal["per-invocation", "per-thread"],
|
||||
) -> None:
|
||||
app = _nested_app(
|
||||
sync_checkpointer,
|
||||
pause_before_b=True,
|
||||
subgraph_checkpointer=True if persistence == "per-thread" else None,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
(task,) = app.get_state(config, subgraphs=True).tasks
|
||||
|
||||
assert task.state.values == _both("a1")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
|
||||
async def test_interrupted_subgraph_task_state_async(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
persistence: Literal["per-invocation", "per-thread"],
|
||||
) -> None:
|
||||
app = _nested_app(
|
||||
async_checkpointer,
|
||||
pause_before_b=True,
|
||||
subgraph_checkpointer=True if persistence == "per-thread" else None,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await app.ainvoke({}, config)
|
||||
|
||||
(task,) = (await app.aget_state(config, subgraphs=True)).tasks
|
||||
|
||||
assert task.state.values == _both("a1")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
|
||||
def test_interrupted_subgraph_history_from_its_task_config(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
persistence: Literal["per-invocation", "per-thread"],
|
||||
) -> None:
|
||||
app = _nested_app(
|
||||
sync_checkpointer,
|
||||
pause_before_b=True,
|
||||
subgraph_checkpointer=True if persistence == "per-thread" else None,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
(task,) = app.get_state(config).tasks
|
||||
|
||||
history = list(app.get_state_history(task.state))
|
||||
|
||||
assert [snapshot.values for snapshot in history][:1] == [_both("a1")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
|
||||
async def test_interrupted_subgraph_ahistory_from_its_task_config(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
persistence: Literal["per-invocation", "per-thread"],
|
||||
) -> None:
|
||||
app = _nested_app(
|
||||
async_checkpointer,
|
||||
pause_before_b=True,
|
||||
subgraph_checkpointer=True if persistence == "per-thread" else None,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await app.ainvoke({}, config)
|
||||
(task,) = (await app.aget_state(config)).tasks
|
||||
|
||||
history = [snapshot async for snapshot in app.aget_state_history(task.state)]
|
||||
|
||||
assert [snapshot.values for snapshot in history][:1] == [_both("a1")]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
|
||||
def test_interrupted_subgraph_update_state_from_its_task_config(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
persistence: Literal["per-invocation", "per-thread"],
|
||||
) -> None:
|
||||
app = _nested_app(
|
||||
sync_checkpointer,
|
||||
pause_before_b=True,
|
||||
subgraph_checkpointer=True if persistence == "per-thread" else None,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
(task,) = app.get_state(config).tasks
|
||||
|
||||
app.update_state(task.state, _both("edit"), as_node="a")
|
||||
|
||||
(task,) = app.get_state(config, subgraphs=True).tasks
|
||||
assert task.state.values == _both("a1", "edit")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("persistence", ["per-invocation", "per-thread"])
|
||||
async def test_interrupted_subgraph_aupdate_state_from_its_task_config(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
persistence: Literal["per-invocation", "per-thread"],
|
||||
) -> None:
|
||||
app = _nested_app(
|
||||
async_checkpointer,
|
||||
pause_before_b=True,
|
||||
subgraph_checkpointer=True if persistence == "per-thread" else None,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await app.ainvoke({}, config)
|
||||
(task,) = (await app.aget_state(config)).tasks
|
||||
|
||||
await app.aupdate_state(task.state, _both("edit"), as_node="a")
|
||||
|
||||
(task,) = (await app.aget_state(config, subgraphs=True)).tasks
|
||||
assert task.state.values == _both("a1", "edit")
|
||||
|
||||
|
||||
def test_subgraph_update_state_keeps_history(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
app = _nested_app(sync_checkpointer, snapshot_frequency=2)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
child = _scoped(config, _child_namespace(app, config))
|
||||
app.update_state(child, _both("manual"))
|
||||
|
||||
assert app.get_state(child).values == _both("a1", "b1", "b2", "manual")
|
||||
|
||||
|
||||
async def test_subgraph_aupdate_state_keeps_history(
|
||||
async_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
app = _nested_app(async_checkpointer, snapshot_frequency=2)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
await app.ainvoke({}, config)
|
||||
|
||||
child = _scoped(config, await _achild_namespace(app, config))
|
||||
await app.aupdate_state(child, _both("manual"))
|
||||
|
||||
assert (await app.aget_state(child)).values == _both("a1", "b1", "b2", "manual")
|
||||
|
||||
|
||||
def test_stateless_subgraph_persists_nothing(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
app = _nested_app(sync_checkpointer, subgraph_checkpointer=False)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
child_tasks = [
|
||||
task
|
||||
for snapshot in app.get_state_history(config)
|
||||
for task in snapshot.tasks
|
||||
if task.name == "child" and isinstance(task.state, dict)
|
||||
]
|
||||
|
||||
assert child_tasks == []
|
||||
assert app.get_state(config).values == _both("a1", "b1", "b2")
|
||||
|
||||
|
||||
def test_completed_subgraph_exposes_no_task_state(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
) -> None:
|
||||
app = _nested_app(sync_checkpointer)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
app.invoke({}, config)
|
||||
|
||||
assert app.get_state(config, subgraphs=True).tasks == ()
|
||||
|
||||
|
||||
def _written_delta_checkpoint() -> Any:
|
||||
checkpoint = empty_checkpoint()
|
||||
checkpoint["channel_versions"]["delta"] = 1
|
||||
return checkpoint
|
||||
|
||||
|
||||
def test_hydrating_written_delta_channel_without_saver_raises() -> None:
|
||||
with pytest.raises(ValueError, match="no checkpointer"):
|
||||
channels_from_checkpoint(
|
||||
{"delta": DeltaChannel(_extend)}, _written_delta_checkpoint()
|
||||
)
|
||||
|
||||
|
||||
async def test_ahydrating_written_delta_channel_without_saver_raises() -> None:
|
||||
with pytest.raises(ValueError, match="no checkpointer"):
|
||||
await achannels_from_checkpoint(
|
||||
{"delta": DeltaChannel(_extend)}, _written_delta_checkpoint()
|
||||
)
|
||||
|
||||
|
||||
def test_hydrating_written_delta_channel_without_config_raises() -> None:
|
||||
with pytest.raises(ValueError, match="no checkpointer"):
|
||||
channels_from_checkpoint(
|
||||
{"delta": DeltaChannel(_extend)},
|
||||
_written_delta_checkpoint(),
|
||||
saver=InMemorySaver(),
|
||||
)
|
||||
|
||||
|
||||
def test_hydrating_unwritten_delta_channel_without_saver_is_empty() -> None:
|
||||
channels, _ = channels_from_checkpoint(
|
||||
{"delta": DeltaChannel(_extend)}, empty_checkpoint()
|
||||
)
|
||||
assert channels["delta"].get() == []
|
||||
|
||||
|
||||
def test_root_checkpointer_true_graph_state_read_raises() -> None:
|
||||
app = _child_builder().compile(checkpointer=True)
|
||||
with pytest.raises(RuntimeError, match="checkpointer=True cannot be used"):
|
||||
app.get_state({"configurable": {"thread_id": "1"}})
|
||||
|
||||
|
||||
def test_stateless_graph_update_state_ignores_lent_saver() -> None:
|
||||
saver = InMemorySaver()
|
||||
app = _child_builder().compile(checkpointer=False)
|
||||
config = {
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"checkpoint_ns": "child:1",
|
||||
CONFIG_KEY_CHECKPOINTER: saver,
|
||||
}
|
||||
}
|
||||
with pytest.raises(ValueError, match="No checkpointer set"):
|
||||
app.update_state(config, _both("x"), as_node="a")
|
||||
assert list(saver.list(None)) == []
|
||||
Reference in New Issue
Block a user