mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 12:35:08 +02:00
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:
@@ -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 []
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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])
|
||||
|
||||
|
||||
Reference in New Issue
Block a user