mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 10:17:50 +02:00
return all checkpoints from .list
This commit is contained in:
@@ -253,9 +253,6 @@ class BasePostgresSaver(BaseCheckpointSaver):
|
||||
if config:
|
||||
wheres.append("thread_id = %s ")
|
||||
param_values.append(config["configurable"]["thread_id"])
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
wheres.append("checkpoint_ns = %s")
|
||||
param_values.append(checkpoint_ns)
|
||||
|
||||
# construct predicate for metadata filter
|
||||
if filter:
|
||||
|
||||
@@ -70,9 +70,6 @@ def search_where(
|
||||
if config is not None:
|
||||
wheres.append("thread_id = ?")
|
||||
param_values.append(config["configurable"]["thread_id"])
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
wheres.append("checkpoint_ns = ?")
|
||||
param_values.append(checkpoint_ns)
|
||||
|
||||
# construct predicate for metadata filter
|
||||
if filter:
|
||||
|
||||
@@ -259,7 +259,6 @@ class BaseCheckpointSaver(ABC):
|
||||
filter: Optional[Dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
include_nested_checkpoints: bool = False,
|
||||
) -> Iterator[CheckpointTuple]:
|
||||
"""List checkpoints that match the given criteria.
|
||||
|
||||
@@ -351,7 +350,6 @@ class BaseCheckpointSaver(ABC):
|
||||
filter: Optional[Dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
include_nested_checkpoints: bool = False,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
"""Asynchronously list checkpoints that match the given criteria.
|
||||
|
||||
|
||||
@@ -157,7 +157,6 @@ class MemorySaver(
|
||||
filter: Optional[Dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
include_nested_checkpoints: bool = False,
|
||||
) -> Iterator[CheckpointTuple]:
|
||||
"""List checkpoints from the in-memory storage.
|
||||
|
||||
@@ -178,16 +177,7 @@ class MemorySaver(
|
||||
config["configurable"].get("checkpoint_ns", "") if config else ""
|
||||
)
|
||||
for thread_id in thread_ids:
|
||||
checkpoint_ns_iter = (
|
||||
(
|
||||
key
|
||||
for key in self.storage[thread_id].keys()
|
||||
if key.startswith(checkpoint_ns)
|
||||
)
|
||||
if include_nested_checkpoints
|
||||
else [checkpoint_ns]
|
||||
)
|
||||
for checkpoint_ns in checkpoint_ns_iter:
|
||||
for checkpoint_ns in self.storage[thread_id].keys():
|
||||
for checkpoint_id, (
|
||||
checkpoint,
|
||||
metadata_b,
|
||||
@@ -330,7 +320,6 @@ class MemorySaver(
|
||||
filter: Optional[Dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
include_nested_checkpoints: bool = False,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
"""Asynchronous version of list.
|
||||
|
||||
@@ -351,7 +340,6 @@ class MemorySaver(
|
||||
before=before,
|
||||
limit=limit,
|
||||
filter=filter,
|
||||
include_nested_checkpoints=include_nested_checkpoints,
|
||||
),
|
||||
config,
|
||||
)
|
||||
|
||||
@@ -360,70 +360,64 @@ class Pregel(
|
||||
if is_managed_value(v)
|
||||
}
|
||||
|
||||
def _prepare_state_snapshot(
|
||||
self, saved: CheckpointTuple, config: RunnableConfig
|
||||
) -> StateSnapshot:
|
||||
checkpoint = saved.checkpoint if saved else empty_checkpoint()
|
||||
config = saved.config if saved else config
|
||||
def _prepare_state_snapshot(self, saved: CheckpointTuple) -> StateSnapshot:
|
||||
with ChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in self.channels.items()
|
||||
},
|
||||
checkpoint,
|
||||
config,
|
||||
saved.checkpoint,
|
||||
saved.config,
|
||||
) as channels, ManagedValuesManager(
|
||||
self.managed_values_dict, ensure_config(config)
|
||||
self.managed_values_dict, ensure_config(saved.config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
checkpoint,
|
||||
saved.checkpoint,
|
||||
self.nodes,
|
||||
channels,
|
||||
managed,
|
||||
config,
|
||||
saved.config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
return StateSnapshot(
|
||||
values=read_channels(channels, self.stream_channels_asis),
|
||||
next=tuple(t.name for t in next_tasks),
|
||||
config=saved.config if saved else config,
|
||||
metadata=saved.metadata if saved else None,
|
||||
created_at=saved.checkpoint["ts"] if saved else None,
|
||||
parent_config=saved.parent_config if saved else None,
|
||||
config=saved.config,
|
||||
metadata=saved.metadata,
|
||||
created_at=saved.checkpoint["ts"],
|
||||
parent_config=saved.parent_config,
|
||||
)
|
||||
|
||||
async def _prepare_state_snapshot_async(
|
||||
self, saved: CheckpointTuple, config: RunnableConfig
|
||||
self, saved: CheckpointTuple
|
||||
) -> StateSnapshot:
|
||||
checkpoint = saved.checkpoint if saved else empty_checkpoint()
|
||||
config = saved.config if saved else config
|
||||
async with AsyncChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in self.channels.items()
|
||||
},
|
||||
checkpoint,
|
||||
config,
|
||||
saved.checkpoint,
|
||||
saved.config,
|
||||
) as channels, AsyncManagedValuesManager(
|
||||
self.managed_values_dict, ensure_config(config)
|
||||
self.managed_values_dict, ensure_config(saved.config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
checkpoint,
|
||||
saved.checkpoint,
|
||||
self.nodes,
|
||||
channels,
|
||||
managed,
|
||||
config,
|
||||
saved.config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
return StateSnapshot(
|
||||
values=read_channels(channels, self.stream_channels_asis),
|
||||
next=tuple(t.name for t in next_tasks),
|
||||
config=saved.config if saved else config,
|
||||
metadata=saved.metadata if saved else None,
|
||||
created_at=saved.checkpoint["ts"] if saved else None,
|
||||
parent_config=saved.parent_config if saved else None,
|
||||
config=saved.config,
|
||||
metadata=saved.metadata,
|
||||
created_at=saved.checkpoint["ts"],
|
||||
parent_config=saved.parent_config,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -457,7 +451,9 @@ class Pregel(
|
||||
|
||||
state_snapshot = checkpoint_ns_to_state_snapshots.pop(root_checkpoint_ns, None)
|
||||
if state_snapshot is None:
|
||||
raise ValueError(f"Missing checkpoint for thread ID '{root_checkpoint_ns}'")
|
||||
raise ValueError(
|
||||
f"Missing checkpoint for checkpoint NS '{root_checkpoint_ns}'"
|
||||
)
|
||||
return state_snapshot
|
||||
|
||||
def get_state(
|
||||
@@ -468,9 +464,7 @@ class Pregel(
|
||||
raise ValueError("No checkpointer set")
|
||||
|
||||
if include_subgraph_state:
|
||||
checkpoint_tuples = self.checkpointer.list(
|
||||
config, include_nested_checkpoints=True
|
||||
)
|
||||
checkpoint_tuples = self.checkpointer.list(config)
|
||||
else:
|
||||
checkpoint_tuples = iter([self.checkpointer.get_tuple(config)])
|
||||
|
||||
@@ -496,20 +490,16 @@ class Pregel(
|
||||
existing_checkpoint_id is None
|
||||
or saved_checkpoint_id > existing_checkpoint_id
|
||||
):
|
||||
state_snapshot = self._prepare_state_snapshot(checkpoint_tuple, config)
|
||||
state_snapshot = self._prepare_state_snapshot(checkpoint_tuple)
|
||||
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
|
||||
checkpoint_ns_to_checkpoint_id[
|
||||
saved_checkpoint_ns
|
||||
] = saved_checkpoint_id
|
||||
|
||||
if not checkpoint_ns_to_state_snapshots:
|
||||
error_msg = (
|
||||
f"Could not find checkpoints for checkpoint NS '{checkpoint_ns}'"
|
||||
return StateSnapshot(
|
||||
values={}, next=(), config=config, checkpoint=empty_checkpoint()
|
||||
)
|
||||
if checkpoint_id:
|
||||
error_msg += f" and checkpoint ID '{checkpoint_id}'"
|
||||
|
||||
raise ValueError(error_msg)
|
||||
|
||||
state_snapshot = self._assemble_state_snapshot_hierarchy(
|
||||
checkpoint_ns, checkpoint_ns_to_state_snapshots
|
||||
@@ -524,9 +514,7 @@ class Pregel(
|
||||
raise ValueError("No checkpointer set")
|
||||
|
||||
if include_subgraph_state:
|
||||
checkpoint_tuples = self.checkpointer.alist(
|
||||
config, include_nested_checkpoints=True
|
||||
)
|
||||
checkpoint_tuples = self.checkpointer.alist(config)
|
||||
else:
|
||||
|
||||
async def alist_checkpoints():
|
||||
@@ -556,20 +544,18 @@ class Pregel(
|
||||
existing_checkpoint_id is None
|
||||
or saved_checkpoint_id > existing_checkpoint_id
|
||||
):
|
||||
state_snapshot = self._prepare_state_snapshot(checkpoint_tuple, config)
|
||||
state_snapshot = await self._prepare_state_snapshot_async(
|
||||
checkpoint_tuple
|
||||
)
|
||||
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
|
||||
checkpoint_ns_to_checkpoint_id[
|
||||
saved_checkpoint_ns
|
||||
] = saved_checkpoint_id
|
||||
|
||||
if not checkpoint_ns_to_state_snapshots:
|
||||
error_msg = (
|
||||
f"Could not find checkpoints for checkpoint NS '{checkpoint_ns}'"
|
||||
return StateSnapshot(
|
||||
values={}, next=(), config=config, checkpoint=empty_checkpoint()
|
||||
)
|
||||
if checkpoint_id:
|
||||
error_msg += f" and checkpoint ID '{checkpoint_id}'"
|
||||
|
||||
raise ValueError(error_msg)
|
||||
|
||||
state_snapshot = self._assemble_state_snapshot_hierarchy(
|
||||
checkpoint_ns, checkpoint_ns_to_state_snapshots
|
||||
@@ -593,41 +579,25 @@ class Pregel(
|
||||
and signature(self.checkpointer.list).parameters.get("filter") is None
|
||||
):
|
||||
raise ValueError("Checkpointer does not support filtering")
|
||||
for config, checkpoint, metadata, parent_config, _ in self.checkpointer.list(
|
||||
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
for checkpoint_tuple in self.checkpointer.list(
|
||||
config, before=before, limit=limit, filter=filter
|
||||
):
|
||||
if (
|
||||
checkpoint_tuple.config["configurable"]["checkpoint_ns"]
|
||||
!= checkpoint_ns
|
||||
):
|
||||
# only list root checkpoints here
|
||||
continue
|
||||
|
||||
if include_subgraph_state:
|
||||
state_snapshot = self.get_state(config, include_subgraph_state=True)
|
||||
state_snapshot = self.get_state(
|
||||
checkpoint_tuple.config, include_subgraph_state=True
|
||||
)
|
||||
yield state_snapshot
|
||||
else:
|
||||
with ChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in self.channels.items()
|
||||
},
|
||||
checkpoint,
|
||||
config,
|
||||
) as channels, ManagedValuesManager(
|
||||
self.managed_values_dict, ensure_config(config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
checkpoint,
|
||||
self.nodes,
|
||||
channels,
|
||||
managed,
|
||||
config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
|
||||
yield StateSnapshot(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(t.name for t in next_tasks),
|
||||
config,
|
||||
metadata,
|
||||
checkpoint["ts"],
|
||||
parent_config,
|
||||
)
|
||||
yield self._prepare_state_snapshot(checkpoint_tuple)
|
||||
|
||||
async def aget_state_history(
|
||||
self,
|
||||
@@ -646,46 +616,25 @@ class Pregel(
|
||||
and signature(self.checkpointer.list).parameters.get("filter") is None
|
||||
):
|
||||
raise ValueError("Checkpointer does not support filtering")
|
||||
async for (
|
||||
config,
|
||||
checkpoint,
|
||||
metadata,
|
||||
parent_config,
|
||||
_,
|
||||
) in self.checkpointer.alist(config, before=before, limit=limit, filter=filter):
|
||||
|
||||
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
||||
async for checkpoint_tuple in self.checkpointer.alist(
|
||||
config, before=before, limit=limit, filter=filter
|
||||
):
|
||||
if (
|
||||
checkpoint_tuple.config["configurable"]["checkpoint_ns"]
|
||||
!= checkpoint_ns
|
||||
):
|
||||
# only list root checkpoints here
|
||||
continue
|
||||
|
||||
if include_subgraph_state:
|
||||
state_snapshot = await self.aget_state(
|
||||
config, include_subgraph_state=True
|
||||
checkpoint_tuple.config, include_subgraph_state=True
|
||||
)
|
||||
yield state_snapshot
|
||||
else:
|
||||
async with AsyncChannelsManager(
|
||||
{
|
||||
k: LastValue(None) if isinstance(c, Context) else c
|
||||
for k, c in self.channels.items()
|
||||
},
|
||||
checkpoint,
|
||||
config,
|
||||
) as channels, AsyncManagedValuesManager(
|
||||
self.managed_values_dict, ensure_config(config)
|
||||
) as managed:
|
||||
next_tasks = prepare_next_tasks(
|
||||
checkpoint,
|
||||
self.nodes,
|
||||
channels,
|
||||
managed,
|
||||
config,
|
||||
-1,
|
||||
for_execution=False,
|
||||
)
|
||||
yield StateSnapshot(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(t.name for t in next_tasks),
|
||||
config,
|
||||
metadata,
|
||||
checkpoint["ts"],
|
||||
parent_config,
|
||||
)
|
||||
yield await self._prepare_state_snapshot_async(checkpoint_tuple)
|
||||
|
||||
def update_state(
|
||||
self,
|
||||
@@ -953,32 +902,36 @@ class Pregel(
|
||||
stream_mode = stream_mode if stream_mode is not None else self.stream_mode
|
||||
if not isinstance(stream_mode, list):
|
||||
stream_mode = [stream_mode]
|
||||
|
||||
if config and config.get("configurable", {}).get(CONFIG_KEY_CHECKPOINTER):
|
||||
parent_checkpointer = config["configurable"][CONFIG_KEY_CHECKPOINTER]
|
||||
else:
|
||||
parent_checkpointer = None
|
||||
|
||||
if config and config.get("configurable", {}).get(CONFIG_KEY_READ) is not None:
|
||||
# if being called as a node in another graph, always use values mode
|
||||
stream_mode = ["values"]
|
||||
|
||||
if self.checkpointer is None:
|
||||
if parent_checkpointer is not None and self.checkpointer is None:
|
||||
raise ValueError(
|
||||
"Missing checkpointer for subgraph. "
|
||||
"Please compile the subgraph graph with checkpointer=INHERIT_CHECKPOINTER (from langgraph.pregel import INHERIT_CHECKPOINTER)."
|
||||
)
|
||||
|
||||
if self.checkpointer != INHERIT_CHECKPOINTER:
|
||||
if (
|
||||
parent_checkpointer is not None
|
||||
and self.checkpointer != INHERIT_CHECKPOINTER
|
||||
):
|
||||
raise ValueError(
|
||||
"Custom checkpointers for subgraphs are not allowed. "
|
||||
"Please compile the subgraph graph with checkpointer=INHERIT_CHECKPOINTER (from langgraph.pregel import INHERIT_CHECKPOINTER)."
|
||||
)
|
||||
|
||||
if (
|
||||
config is not None
|
||||
and config.get("configurable", {}).get(CONFIG_KEY_CHECKPOINTER)
|
||||
and self.checkpointer == INHERIT_CHECKPOINTER
|
||||
):
|
||||
checkpointer: Optional[BaseCheckpointSaver] = config["configurable"][
|
||||
CONFIG_KEY_CHECKPOINTER
|
||||
]
|
||||
else:
|
||||
checkpointer = self.checkpointer
|
||||
checkpointer = (
|
||||
parent_checkpointer
|
||||
if parent_checkpointer is not None
|
||||
else self.checkpointer
|
||||
)
|
||||
return (
|
||||
debug,
|
||||
stream_mode,
|
||||
|
||||
@@ -7629,13 +7629,13 @@ async def test_nested_graph_state(
|
||||
my_key: str
|
||||
my_other_key: str
|
||||
|
||||
def inner_1(state: InnerState):
|
||||
async def inner_1(state: InnerState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " here",
|
||||
"my_other_key": state["my_key"],
|
||||
}
|
||||
|
||||
def inner_2(state: InnerState):
|
||||
async def inner_2(state: InnerState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " and there",
|
||||
"my_other_key": state["my_key"],
|
||||
@@ -7651,10 +7651,10 @@ async def test_nested_graph_state(
|
||||
class State(TypedDict):
|
||||
my_key: str
|
||||
|
||||
def outer_1(state: State):
|
||||
async def outer_1(state: State):
|
||||
return {"my_key": "hi " + state["my_key"]}
|
||||
|
||||
def outer_2(state: State):
|
||||
async def outer_2(state: State):
|
||||
return {"my_key": state["my_key"] + " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
@@ -7699,7 +7699,7 @@ async def test_nested_graph_state(
|
||||
},
|
||||
subgraph_state_snapshots=None,
|
||||
)
|
||||
assert app.get_state(config, include_subgraph_state=True) == StateSnapshot(
|
||||
assert await app.aget_state(config, include_subgraph_state=True) == StateSnapshot(
|
||||
values={"my_key": "hi my value"},
|
||||
next=("inner",),
|
||||
config={
|
||||
@@ -7755,7 +7755,9 @@ async def test_nested_graph_state(
|
||||
)
|
||||
},
|
||||
)
|
||||
assert list(app.get_state_history(config, include_subgraph_state=True)) == [
|
||||
assert [
|
||||
s async for s in app.aget_state_history(config, include_subgraph_state=True)
|
||||
] == [
|
||||
StateSnapshot(
|
||||
values={"my_key": "hi my value"},
|
||||
next=("inner",),
|
||||
@@ -8089,10 +8091,10 @@ async def test_doubly_nested_graph_state(
|
||||
class GrandChildState(TypedDict):
|
||||
my_key: str
|
||||
|
||||
def grandchild_1(state: ChildState):
|
||||
async def grandchild_1(state: ChildState):
|
||||
return {"my_key": state["my_key"] + " here"}
|
||||
|
||||
def grandchild_2(state: ChildState):
|
||||
async def grandchild_2(state: ChildState):
|
||||
return {
|
||||
"my_key": state["my_key"] + " and there",
|
||||
}
|
||||
@@ -8114,10 +8116,10 @@ async def test_doubly_nested_graph_state(
|
||||
child.set_entry_point("child_1")
|
||||
child.set_finish_point("child_1")
|
||||
|
||||
def parent_1(state: State):
|
||||
async def parent_1(state: State):
|
||||
return {"my_key": "hi " + state["my_key"]}
|
||||
|
||||
def parent_2(state: State):
|
||||
async def parent_2(state: State):
|
||||
return {"my_key": state["my_key"] + " and back again"}
|
||||
|
||||
graph = StateGraph(State)
|
||||
|
||||
Reference in New Issue
Block a user