mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
fix(langgraph): share one saver resolution between runs and state methods
`_defaults` picked the run's saver on its own, and disagreed with the state methods on `checkpointer=False`: a run ignored a lent saver while `update_state` wrote a checkpoint with it that no run reads. A `checkpointer=True` root also got `True` back from `_state_checkpointer` and failed with an AttributeError instead of the run's error. Both now go through `_resolve_checkpointer`, and the state methods require an actual saver.
This commit is contained in:
@@ -236,7 +236,7 @@ def _require_saver_for_history(
|
||||
if written and (saver is None or config is None):
|
||||
raise ValueError(
|
||||
f"DeltaChannel {written} has history to replay but no checkpointer "
|
||||
"and config were passed to read it"
|
||||
"or config was passed to read it"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -835,15 +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
|
||||
)
|
||||
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)
|
||||
if not 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
|
||||
|
||||
@@ -1648,7 +1661,7 @@ class Pregel(
|
||||
channels, managed = channels_from_checkpoint(
|
||||
self.channels,
|
||||
checkpoint,
|
||||
saver=checkpointer if saved is not None else None,
|
||||
saver=checkpointer,
|
||||
config=saved.config if saved is not None else None,
|
||||
)
|
||||
values, as_node = updates[0][:2]
|
||||
@@ -2107,7 +2120,7 @@ class Pregel(
|
||||
channels, managed = await achannels_from_checkpoint(
|
||||
self.channels,
|
||||
checkpoint,
|
||||
saver=checkpointer if saved is not None else None,
|
||||
saver=checkpointer,
|
||||
config=saved.config if saved is not None else None,
|
||||
)
|
||||
values, as_node = updates[0][:2]
|
||||
@@ -2548,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' "
|
||||
|
||||
@@ -6,6 +6,7 @@ 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 (
|
||||
@@ -371,3 +372,24 @@ def test_hydrating_unwritten_delta_channel_without_saver_is_empty() -> None:
|
||||
{"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