mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
fix(langgraph): resolve the saver and checkpointer=True namespace in one place
The six state methods each resolved the saver themselves, and only get_state recast a checkpointer=True subgraph's namespace. So get_state_history and update_state given the subgraph task's config looked under `name:<task_id>`, where the run never stores anything: the history came back empty and the update was silently lost. The state methods and the run loop now share one saver lookup and one namespace rule. The hydration guard also raises when a saver comes without a config, which skipped the replay the same way.
This commit is contained in:
@@ -230,12 +230,13 @@ 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:
|
||||
if written and (saver is None or config is None):
|
||||
raise ValueError(
|
||||
f"DeltaChannel {written} has history to replay but no checkpointer "
|
||||
"was passed to read it"
|
||||
"and config were passed to read it"
|
||||
)
|
||||
|
||||
|
||||
@@ -269,7 +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)
|
||||
_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(
|
||||
@@ -312,7 +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)
|
||||
_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,28 @@ class Pregel(
|
||||
if auto_validate:
|
||||
self.validate()
|
||||
|
||||
def _state_checkpointer(self, config: RunnableConfig) -> BaseCheckpointSaver:
|
||||
"""The checkpointer state reads and writes use: the one a parent lends a
|
||||
subgraph through the config, else this graph's own."""
|
||||
checkpointer = 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")
|
||||
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:
|
||||
@@ -1385,13 +1407,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, "")
|
||||
@@ -1408,11 +1424,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)
|
||||
@@ -1430,13 +1442,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, "")
|
||||
@@ -1453,11 +1459,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)
|
||||
@@ -1481,13 +1483,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, "")
|
||||
@@ -1516,6 +1512,7 @@ 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)
|
||||
@@ -1534,13 +1531,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, "")
|
||||
@@ -1570,6 +1561,7 @@ class Pregel(
|
||||
}
|
||||
},
|
||||
)
|
||||
config = self._own_checkpoint_config(config)
|
||||
# eagerly consume list() to avoid holding up the db cursor
|
||||
for checkpoint_tuple in [
|
||||
c
|
||||
@@ -1602,13 +1594,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")
|
||||
@@ -1635,7 +1621,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)
|
||||
@@ -2065,13 +2053,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")
|
||||
@@ -2098,7 +2080,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)
|
||||
@@ -2792,9 +2776,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))
|
||||
@@ -3219,9 +3201,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?
|
||||
|
||||
@@ -3,6 +3,7 @@ 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.channels.delta import DeltaChannel
|
||||
@@ -204,6 +205,84 @@ async def test_interrupted_subgraph_task_state_async(
|
||||
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:
|
||||
@@ -278,6 +357,15 @@ async def test_ahydrating_written_delta_channel_without_saver_raises() -> None:
|
||||
)
|
||||
|
||||
|
||||
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()
|
||||
|
||||
Reference in New Issue
Block a user