From 87cb5095285ba809b547584d31bae26046fc844a Mon Sep 17 00:00:00 2001 From: Imvikram99 <120546792+Imvikram99@users.noreply.github.com> Date: Mon, 15 Dec 2025 12:52:47 +0530 Subject: [PATCH] feat(checkpoint): Validate checkpointer type at compile time (#6586) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Description: Catch invalid checkpointer objects early by validating any checkpointer argument before compilation/execution, raising a clear TypeError that instructs users to pass a proper BaseCheckpointSaver (e.g., AsyncPostgresSaver) instead of stores like AsyncPostgresStore. Includes shared validation logic and a regression test so we don’t see AttributeError: 'AsyncPostgresStore' object has no attribute 'get_next_version' again. Issue: Fixes #6585 Dependencies: None Twitter handle: none --- libs/langgraph/langgraph/graph/state.py | 3 +++ libs/langgraph/langgraph/pregel/main.py | 5 ++++- libs/langgraph/langgraph/types.py | 15 +++++++++++++++ libs/langgraph/tests/test_pregel.py | 16 ++++++++++++++++ libs/prebuilt/tests/test_react_agent.py | 2 +- 5 files changed, 39 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 2654ee057..f3299de92 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -78,6 +78,7 @@ from langgraph.types import ( Command, RetryPolicy, Send, + ensure_valid_checkpointer, ) from langgraph.typing import ContextT, InputT, NodeInputT, OutputT, StateT from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV10 @@ -853,6 +854,8 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]): Returns: CompiledStateGraph: The compiled `StateGraph`. """ + checkpointer = ensure_valid_checkpointer(checkpointer) + # assign default values interrupt_before = interrupt_before or [] interrupt_after = interrupt_after or [] diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index a59becfa0..37e8125f9 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -142,6 +142,7 @@ from langgraph.types import ( StateSnapshot, StateUpdate, StreamMode, + ensure_valid_checkpointer, ) from langgraph.typing import ContextT, InputT, OutputT, StateT from langgraph.warnings import LangGraphDeprecatedSinceV10 @@ -642,7 +643,7 @@ class Pregel( input_channels: str | Sequence[str], step_timeout: float | None = None, debug: bool | None = None, - checkpointer: BaseCheckpointSaver | None = None, + checkpointer: Checkpointer = None, store: BaseStore | None = None, cache: BaseCache | None = None, retry_policy: RetryPolicy | Sequence[RetryPolicy] = (), @@ -665,6 +666,8 @@ class Pregel( if context_schema is None: context_schema = cast(type[ContextT], config_type) + checkpointer = ensure_valid_checkpointer(checkpointer) + self.nodes = { k: v.build() if isinstance(v, NodeBuilder) else v for k, v in nodes.items() } diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 18e8854c7..c8a798a38 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -56,6 +56,7 @@ __all__ = ( "Durability", "interrupt", "Overwrite", + "ensure_valid_checkpointer", ) Durability = Literal["sync", "async", "exit"] @@ -73,6 +74,20 @@ Checkpointer = None | bool | BaseCheckpointSaver - False disables checkpointing, even if the parent graph has a checkpointer. - None inherits checkpointer from the parent graph.""" + +def ensure_valid_checkpointer(checkpointer: Checkpointer) -> Checkpointer: + if checkpointer not in (None, True, False) and not isinstance( + checkpointer, BaseCheckpointSaver + ): + raise TypeError( + "Invalid checkpointer provided. Expected an instance of " + "`BaseCheckpointSaver`, `True`, `False`, or `None`. " + f"Received {type(checkpointer).__name__!s}. " + "Pass a proper saver (e.g., InMemorySaver, AsyncPostgresSaver)." + ) + return checkpointer + + StreamMode = Literal[ "values", "updates", "checkpoints", "tasks", "debug", "messages", "custom" ] diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 5ede65afc..97928ec6f 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -120,6 +120,22 @@ def test_graph_validation() -> None: graph.invoke({"hello": "there"}) +def test_invalid_checkpointer_type() -> None: + class State(TypedDict): + foo: str + + builder = StateGraph(State) + builder.add_node("start", lambda state: state) + builder.set_entry_point("start") + builder.set_finish_point("start") + + class NotACheckpointer: + pass + + with pytest.raises(TypeError, match="Invalid checkpointer provided"): + builder.compile(checkpointer=NotACheckpointer()) + + def test_graph_validation_with_command() -> None: class State(TypedDict): foo: str diff --git a/libs/prebuilt/tests/test_react_agent.py b/libs/prebuilt/tests/test_react_agent.py index 4dd265d4b..c801a96dc 100644 --- a/libs/prebuilt/tests/test_react_agent.py +++ b/libs/prebuilt/tests/test_react_agent.py @@ -638,7 +638,7 @@ def test_react_agent_parallel_tool_calls( for event in agent.stream( {"messages": [("user", query)]}, config, stream_mode="values" ): - if "__interrupt__" not in event: + if "__interrupt__" not in event: if messages := event.get("messages"): message_types.append([m.type for m in messages])