feat(checkpoint): Validate checkpointer type at compile time (#6586)

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
This commit is contained in:
Imvikram99
2025-12-15 07:22:47 +00:00
committed by GitHub
parent 84023451a2
commit 87cb509528
5 changed files with 39 additions and 2 deletions
+3
View File
@@ -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 []
+4 -1
View File
@@ -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()
}
+15
View File
@@ -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"
]
+16
View File
@@ -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
+1 -1
View File
@@ -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])