mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 04:55:09 +02:00
Add checkpoint metadata fields
- source: input, update or loop - step: int - make step counter continue from previous last step
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
import asyncio
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from types import TracebackType
|
||||
from typing import Any, AsyncIterator, Optional
|
||||
from typing import AsyncIterator, Optional
|
||||
|
||||
import aiosqlite
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
@@ -10,6 +10,7 @@ from typing_extensions import Self
|
||||
from langgraph.checkpoint.base import (
|
||||
BaseCheckpointSaver,
|
||||
Checkpoint,
|
||||
CheckpointMetadata,
|
||||
CheckpointTuple,
|
||||
SerializerProtocol,
|
||||
)
|
||||
@@ -164,7 +165,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
return CheckpointTuple(
|
||||
config,
|
||||
self.serde.loads(value[0]),
|
||||
self.serde.loads(value[2]) if value[2] is not None else None,
|
||||
self.serde.loads(value[2]) if value[2] is not None else {},
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": config["configurable"]["thread_id"],
|
||||
@@ -188,7 +189,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
}
|
||||
},
|
||||
self.serde.loads(value[3]),
|
||||
self.serde.loads(value[4]) if value[4] is not None else None,
|
||||
self.serde.loads(value[4]) if value[4] is not None else {},
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": value[0],
|
||||
@@ -242,7 +243,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
yield CheckpointTuple(
|
||||
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
|
||||
self.serde.loads(value),
|
||||
self.serde.loads(metadata) if metadata is not None else None,
|
||||
self.serde.loads(metadata) if metadata is not None else {},
|
||||
{"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}}
|
||||
if parent_ts
|
||||
else None,
|
||||
@@ -252,7 +253,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: Optional[dict[str, Any]] = None,
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the database asynchronously.
|
||||
|
||||
@@ -274,7 +275,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
||||
checkpoint["ts"],
|
||||
config["configurable"].get("thread_ts"),
|
||||
self.serde.dumps(checkpoint),
|
||||
self.serde.dumps(metadata) if metadata is not None else None,
|
||||
self.serde.dumps(metadata),
|
||||
),
|
||||
):
|
||||
await self.conn.commit()
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import (
|
||||
Any,
|
||||
AsyncIterator,
|
||||
Iterator,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
TypedDict,
|
||||
@@ -16,6 +17,22 @@ from langgraph.serde.base import SerializerProtocol
|
||||
from langgraph.serde.jsonplus import JsonPlusSerializer
|
||||
|
||||
|
||||
# Marked as total=False to allow for future expansion.
|
||||
class CheckpointMetadata(TypedDict, total=False):
|
||||
source: Literal["input", "loop", "update"]
|
||||
"""The source of the checkpoint.
|
||||
- "input": The checkpoint was created from an input to invoke/stream/batch.
|
||||
- "loop": The checkpoint was created from inside the pregel loop.
|
||||
- "update": The checkpoint was created from a manual state update.
|
||||
"""
|
||||
step: int
|
||||
"""The step number of the checkpoint.
|
||||
-1 for the first "input" checkpoint.
|
||||
0 for the first "loop" checkpoint.
|
||||
... for the nth checkpoint afterwards.
|
||||
"""
|
||||
|
||||
|
||||
class Checkpoint(TypedDict):
|
||||
"""State snapshot at a given point in time."""
|
||||
|
||||
@@ -73,7 +90,7 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
|
||||
class CheckpointTuple(NamedTuple):
|
||||
config: RunnableConfig
|
||||
checkpoint: Checkpoint
|
||||
metadata: Optional[dict[str, Any]]
|
||||
metadata: CheckpointMetadata
|
||||
parent_config: Optional[RunnableConfig] = None
|
||||
|
||||
|
||||
@@ -130,7 +147,7 @@ class BaseCheckpointSaver(ABC):
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: dict[str, Any],
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -154,6 +171,6 @@ class BaseCheckpointSaver(ABC):
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: dict[str, Any],
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
import asyncio
|
||||
from collections import defaultdict
|
||||
from typing import Any, AsyncIterator, Iterator, Optional
|
||||
from typing import AsyncIterator, Iterator, Optional
|
||||
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
from langgraph.checkpoint.base import (
|
||||
BaseCheckpointSaver,
|
||||
Checkpoint,
|
||||
CheckpointMetadata,
|
||||
CheckpointTuple,
|
||||
SerializerProtocol,
|
||||
)
|
||||
@@ -118,7 +119,7 @@ class MemorySaver(BaseCheckpointSaver):
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: dict[str, Any] = None,
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the in-memory storage.
|
||||
|
||||
@@ -136,7 +137,7 @@ class MemorySaver(BaseCheckpointSaver):
|
||||
{
|
||||
checkpoint["ts"]: (
|
||||
self.serde.dumps(checkpoint),
|
||||
self.serde.dumps(metadata or {}),
|
||||
self.serde.dumps(metadata),
|
||||
)
|
||||
}
|
||||
)
|
||||
@@ -184,8 +185,11 @@ class MemorySaver(BaseCheckpointSaver):
|
||||
return
|
||||
|
||||
async def aput(
|
||||
self, config: RunnableConfig, checkpoint: Checkpoint
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
return await asyncio.get_running_loop().run_in_executor(
|
||||
None, self.put, config, checkpoint
|
||||
None, self.put, config, checkpoint, metadata
|
||||
)
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing_extensions import Self
|
||||
from langgraph.checkpoint.base import (
|
||||
BaseCheckpointSaver,
|
||||
Checkpoint,
|
||||
CheckpointMetadata,
|
||||
CheckpointTuple,
|
||||
SerializerProtocol,
|
||||
)
|
||||
@@ -220,7 +221,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
return CheckpointTuple(
|
||||
config,
|
||||
self.serde.loads(value[0]),
|
||||
self.serde.loads(value[2]) if value[2] is not None else None,
|
||||
self.serde.loads(value[2]) if value[2] is not None else {},
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": config["configurable"]["thread_id"],
|
||||
@@ -244,7 +245,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
}
|
||||
},
|
||||
self.serde.loads(value[3]),
|
||||
self.serde.loads(value[4]) if value[4] is not None else None,
|
||||
self.serde.loads(value[4]) if value[4] is not None else {},
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": value[0],
|
||||
@@ -308,7 +309,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
yield CheckpointTuple(
|
||||
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
|
||||
self.serde.loads(value),
|
||||
self.serde.loads(metadata) if metadata is not None else None,
|
||||
self.serde.loads(metadata) if metadata is not None else {},
|
||||
{
|
||||
"configurable": {
|
||||
"thread_id": thread_id,
|
||||
@@ -323,7 +324,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: dict[str, Any] = None,
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
"""Save a checkpoint to the database.
|
||||
|
||||
@@ -355,7 +356,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
||||
checkpoint["ts"],
|
||||
config["configurable"].get("thread_ts"),
|
||||
self.serde.dumps(checkpoint),
|
||||
self.serde.dumps(metadata) if metadata is not None else None,
|
||||
self.serde.dumps(metadata),
|
||||
),
|
||||
)
|
||||
return {
|
||||
|
||||
@@ -342,7 +342,7 @@ class Pregel(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(name for name, _ in next_tasks),
|
||||
config,
|
||||
saved.metadata or {},
|
||||
saved.metadata,
|
||||
)
|
||||
|
||||
async def aget_state(self, config: RunnableConfig) -> StateSnapshot:
|
||||
@@ -361,7 +361,7 @@ class Pregel(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(name for name, _ in next_tasks),
|
||||
config,
|
||||
saved.metadata or {},
|
||||
saved.metadata,
|
||||
)
|
||||
|
||||
def get_state_history(
|
||||
@@ -386,7 +386,7 @@ class Pregel(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(name for name, _ in next_tasks),
|
||||
config,
|
||||
metadata or {},
|
||||
metadata,
|
||||
parent_config,
|
||||
)
|
||||
|
||||
@@ -415,7 +415,7 @@ class Pregel(
|
||||
read_channels(channels, self.stream_channels_asis),
|
||||
tuple(name for name, _ in next_tasks),
|
||||
config,
|
||||
metadata or {},
|
||||
metadata,
|
||||
parent_config,
|
||||
)
|
||||
|
||||
@@ -433,8 +433,8 @@ class Pregel(
|
||||
raise ValueError("No checkpointer set")
|
||||
|
||||
# get last checkpoint
|
||||
checkpoint = self.checkpointer.get(config)
|
||||
checkpoint = copy_checkpoint(checkpoint) if checkpoint else empty_checkpoint()
|
||||
saved = self.checkpointer.get_tuple(config)
|
||||
checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
|
||||
# find last node that updated the state, if not provided
|
||||
if as_node is None:
|
||||
last_seen_by_node = sorted(
|
||||
@@ -482,7 +482,14 @@ class Pregel(
|
||||
# apply to checkpoint and save
|
||||
_apply_writes(checkpoint, channels, task.writes)
|
||||
return self.checkpointer.put(
|
||||
config, create_checkpoint(checkpoint, channels)
|
||||
config,
|
||||
create_checkpoint(checkpoint, channels),
|
||||
{
|
||||
"source": "update",
|
||||
"step": saved.metadata.get("step", 0) + 1
|
||||
if saved.metadata
|
||||
else None,
|
||||
},
|
||||
)
|
||||
|
||||
async def aupdate_state(
|
||||
@@ -495,8 +502,8 @@ class Pregel(
|
||||
raise ValueError("No checkpointer set")
|
||||
|
||||
# get last checkpoint
|
||||
checkpoint = await self.checkpointer.aget(config)
|
||||
checkpoint = copy_checkpoint(checkpoint) if checkpoint else empty_checkpoint()
|
||||
saved = await self.checkpointer.aget_tuple(config)
|
||||
checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
|
||||
# find last node that updated the state, if not provided
|
||||
if as_node is None:
|
||||
last_seen_by_node = sorted(
|
||||
@@ -544,7 +551,12 @@ class Pregel(
|
||||
# apply to checkpoint and save
|
||||
_apply_writes(checkpoint, channels, task.writes)
|
||||
return await self.checkpointer.aput(
|
||||
config, create_checkpoint(checkpoint, channels)
|
||||
config,
|
||||
create_checkpoint(checkpoint, channels),
|
||||
{
|
||||
"source": "update",
|
||||
"step": saved.metadata.get("step", 0) + 1 if saved else None,
|
||||
},
|
||||
)
|
||||
|
||||
def _defaults(
|
||||
@@ -638,10 +650,12 @@ class Pregel(
|
||||
processes = {**self.nodes}
|
||||
# get checkpoint from saver, or create an empty one
|
||||
checkpoint_config = config
|
||||
checkpoint = (
|
||||
self.checkpointer.get(checkpoint_config) if self.checkpointer else None
|
||||
saved = (
|
||||
self.checkpointer.get_tuple(checkpoint_config)
|
||||
if self.checkpointer
|
||||
else None
|
||||
)
|
||||
checkpoint = checkpoint or empty_checkpoint()
|
||||
checkpoint = saved.checkpoint if saved else empty_checkpoint()
|
||||
# create channels from checkpoint
|
||||
with ChannelsManager(
|
||||
self.channels, checkpoint
|
||||
@@ -667,7 +681,9 @@ class Pregel(
|
||||
# channel updates from step N are only visible in step N+1
|
||||
# channels are guaranteed to be immutable for the duration of the step,
|
||||
# with channel updates applied only at the transition between steps
|
||||
for step in range(config["recursion_limit"] + 1):
|
||||
start = saved.metadata.get("step", -1) + 1 if saved else 0
|
||||
stop = start + config["recursion_limit"] + 1
|
||||
for step in range(start, stop):
|
||||
next_checkpoint, next_tasks = _prepare_next_tasks(
|
||||
checkpoint, processes, channels, for_execution=True
|
||||
)
|
||||
@@ -771,7 +787,9 @@ class Pregel(
|
||||
if self.checkpointer is not None:
|
||||
checkpoint = create_checkpoint(checkpoint, channels)
|
||||
checkpoint_config = self.checkpointer.put(
|
||||
checkpoint_config, checkpoint
|
||||
checkpoint_config,
|
||||
checkpoint,
|
||||
{"source": "loop", "step": step},
|
||||
)
|
||||
if stream_mode == "debug":
|
||||
yield map_debug_checkpoint(
|
||||
@@ -865,12 +883,12 @@ class Pregel(
|
||||
processes = {**self.nodes}
|
||||
# get checkpoint from saver, or create an empty one
|
||||
checkpoint_config = config
|
||||
checkpoint = (
|
||||
await self.checkpointer.aget(checkpoint_config)
|
||||
saved = (
|
||||
await self.checkpointer.aget_tuple(checkpoint_config)
|
||||
if self.checkpointer
|
||||
else None
|
||||
)
|
||||
checkpoint = checkpoint or empty_checkpoint()
|
||||
checkpoint = saved.checkpoint if saved else empty_checkpoint()
|
||||
# create channels from checkpoint
|
||||
async with AsyncChannelsManager(self.channels, checkpoint) as channels:
|
||||
# map inputs to channel updates
|
||||
@@ -894,7 +912,9 @@ class Pregel(
|
||||
# channel updates from step N are only visible in step N+1,
|
||||
# channels are guaranteed to be immutable for the duration of the step,
|
||||
# channel updates being applied only at the transition between steps
|
||||
for step in range(config["recursion_limit"] + 1):
|
||||
start = saved.metadata.get("step", -1) + 1 if saved else 0
|
||||
stop = start + config["recursion_limit"] + 1
|
||||
for step in range(start, stop):
|
||||
next_checkpoint, next_tasks = _prepare_next_tasks(
|
||||
checkpoint, processes, channels, for_execution=True
|
||||
)
|
||||
@@ -1008,7 +1028,9 @@ class Pregel(
|
||||
if self.checkpointer is not None:
|
||||
checkpoint = create_checkpoint(checkpoint, channels)
|
||||
checkpoint_config = await self.checkpointer.aput(
|
||||
checkpoint_config, checkpoint
|
||||
checkpoint_config,
|
||||
checkpoint,
|
||||
{"source": "loop", "step": step},
|
||||
)
|
||||
if stream_mode == "debug":
|
||||
yield map_debug_checkpoint(
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Any, Optional
|
||||
|
||||
from langgraph.checkpoint.base import (
|
||||
Checkpoint,
|
||||
CheckpointMetadata,
|
||||
SerializerProtocol,
|
||||
copy_checkpoint,
|
||||
)
|
||||
@@ -30,7 +31,12 @@ class MemorySaverAssertImmutable(MemorySaver):
|
||||
super().__init__(serde=serde)
|
||||
self.storage_for_copies = defaultdict(dict)
|
||||
|
||||
def put(self, config: dict, checkpoint: Checkpoint) -> None:
|
||||
def put(
|
||||
self,
|
||||
config: dict,
|
||||
checkpoint: Checkpoint,
|
||||
metadata: Optional[CheckpointMetadata] = None,
|
||||
) -> None:
|
||||
# assert checkpoint hasn't been modified since last written
|
||||
thread_id = config["configurable"]["thread_id"]
|
||||
if saved := super().get(config):
|
||||
@@ -39,4 +45,4 @@ class MemorySaverAssertImmutable(MemorySaver):
|
||||
checkpoint
|
||||
)
|
||||
# call super to write checkpoint
|
||||
return super().put(config, checkpoint)
|
||||
return super().put(config, checkpoint, metadata)
|
||||
|
||||
+39
-39
@@ -1202,7 +1202,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
assert (
|
||||
app_w_interrupt.checkpointer.get_tuple(config).config["configurable"][
|
||||
@@ -1236,7 +1236,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 1},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -1320,7 +1320,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 4},
|
||||
)
|
||||
|
||||
# test state get/update methods with interrupt_before
|
||||
@@ -1356,7 +1356,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
|
||||
app_w_interrupt.update_state(
|
||||
@@ -1384,7 +1384,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 1},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -1468,7 +1468,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 4},
|
||||
)
|
||||
|
||||
# test re-invoke to continue with interrupt_before
|
||||
@@ -1504,7 +1504,7 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -1836,7 +1836,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
app_w_interrupt.update_state(
|
||||
@@ -1862,7 +1862,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -1921,7 +1921,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
# test state get/update methods with interrupt_before
|
||||
@@ -1956,7 +1956,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
app_w_interrupt.update_state(
|
||||
@@ -1982,7 +1982,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -2041,7 +2041,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
# test w interrupt before all
|
||||
@@ -2064,7 +2064,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("agent",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -2087,7 +2087,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -2126,7 +2126,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("agent",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 2},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -2171,7 +2171,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -2210,7 +2210,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
|
||||
},
|
||||
next=("agent",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 2},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -3025,7 +3025,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
# modify ai message
|
||||
@@ -3051,7 +3051,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("action",),
|
||||
config=next_config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -3117,7 +3117,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 4},
|
||||
)
|
||||
|
||||
app_w_interrupt.update_state(
|
||||
@@ -3153,7 +3153,7 @@ def test_message_graph(
|
||||
],
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
app_w_interrupt = workflow.compile(
|
||||
@@ -3199,7 +3199,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
# modify ai message
|
||||
@@ -3228,7 +3228,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||
@@ -3294,7 +3294,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 4},
|
||||
)
|
||||
|
||||
app_w_interrupt.update_state(
|
||||
@@ -3330,7 +3330,7 @@ def test_message_graph(
|
||||
],
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
# add an extra message as if it came from "action" node
|
||||
@@ -3366,7 +3366,7 @@ def test_message_graph(
|
||||
],
|
||||
next=("agent",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 6},
|
||||
)
|
||||
|
||||
|
||||
@@ -3489,7 +3489,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value", "market": "DE"},
|
||||
next=("tool_two_slow",),
|
||||
config=tool_two.checkpointer.get_tuple(thread1).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert tool_two.invoke(None, thread1, debug=1) == {
|
||||
@@ -3500,7 +3500,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value slow", "market": "DE"},
|
||||
next=(),
|
||||
config=tool_two.checkpointer.get_tuple(thread1).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
@@ -3513,7 +3513,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value", "market": "US"},
|
||||
next=("tool_two_fast",),
|
||||
config=tool_two.checkpointer.get_tuple(thread2).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert tool_two.invoke(None, thread2, debug=1) == {
|
||||
@@ -3524,7 +3524,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value fast", "market": "US"},
|
||||
next=(),
|
||||
config=tool_two.checkpointer.get_tuple(thread2).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
|
||||
@@ -3713,7 +3713,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared", "market": "DE"},
|
||||
next=("tool_two_slow",),
|
||||
config=tool_two.checkpointer.get_tuple(thread1).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert tool_two.invoke(None, thread1, debug=1) == {
|
||||
@@ -3724,7 +3724,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared slow finished", "market": "DE"},
|
||||
next=(),
|
||||
config=tool_two.checkpointer.get_tuple(thread1).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
@@ -3737,7 +3737,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared", "market": "US"},
|
||||
next=("tool_two_fast",),
|
||||
config=tool_two.checkpointer.get_tuple(thread2).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert tool_two.invoke(None, thread2, debug=1) == {
|
||||
@@ -3748,7 +3748,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared fast finished", "market": "US"},
|
||||
next=(),
|
||||
config=tool_two.checkpointer.get_tuple(thread2).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
with SqliteSaver.from_conn_string(":memory:") as saver:
|
||||
@@ -3770,7 +3770,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared", "market": "DE"},
|
||||
next=("tool_two_slow",),
|
||||
config=tool_two.checkpointer.get_tuple(thread1).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert tool_two.invoke(None, thread1, debug=1) == {
|
||||
@@ -3781,7 +3781,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared slow finished", "market": "DE"},
|
||||
next=(),
|
||||
config=tool_two.checkpointer.get_tuple(thread1).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
@@ -3794,7 +3794,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared", "market": "US"},
|
||||
next=("tool_two_fast",),
|
||||
config=tool_two.checkpointer.get_tuple(thread2).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert tool_two.invoke(None, thread2, debug=1) == {
|
||||
@@ -3805,7 +3805,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value prepared fast finished", "market": "US"},
|
||||
next=(),
|
||||
config=tool_two.checkpointer.get_tuple(thread2).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
|
||||
|
||||
+29
-29
@@ -1280,7 +1280,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
|
||||
await app_w_interrupt.aupdate_state(
|
||||
@@ -1308,7 +1308,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 1},
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
@@ -1392,7 +1392,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=(),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 4},
|
||||
)
|
||||
|
||||
# test state get/update methods with interrupt_before
|
||||
@@ -1431,7 +1431,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
|
||||
await app_w_interrupt.aupdate_state(
|
||||
@@ -1459,7 +1459,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 1},
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
@@ -1543,7 +1543,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=(),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 4},
|
||||
)
|
||||
|
||||
# test re-invoke to continue with interrupt_before
|
||||
@@ -1582,7 +1582,7 @@ async def test_conditional_graph() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
@@ -1902,7 +1902,7 @@ async def test_conditional_graph_state() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
await app_w_interrupt.aupdate_state(
|
||||
@@ -1928,7 +1928,7 @@ async def test_conditional_graph_state() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
@@ -1987,7 +1987,7 @@ async def test_conditional_graph_state() -> None:
|
||||
},
|
||||
next=(),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
# test state get/update methods with interrupt_before
|
||||
@@ -2024,7 +2024,7 @@ async def test_conditional_graph_state() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
await app_w_interrupt.aupdate_state(
|
||||
@@ -2050,7 +2050,7 @@ async def test_conditional_graph_state() -> None:
|
||||
},
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
@@ -2109,7 +2109,7 @@ async def test_conditional_graph_state() -> None:
|
||||
},
|
||||
next=(),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
|
||||
@@ -2715,7 +2715,7 @@ async def test_message_graph() -> None:
|
||||
],
|
||||
next=("action",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
# modify ai message
|
||||
@@ -2743,7 +2743,7 @@ async def test_message_graph() -> None:
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 2},
|
||||
)
|
||||
|
||||
assert [c async for c in app_w_interrupt.astream(None, config)] == [
|
||||
@@ -2796,7 +2796,7 @@ async def test_message_graph() -> None:
|
||||
],
|
||||
next=("action",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 4},
|
||||
)
|
||||
|
||||
await app_w_interrupt.aupdate_state(
|
||||
@@ -2830,7 +2830,7 @@ async def test_message_graph() -> None:
|
||||
],
|
||||
next=(),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
metadata={},
|
||||
metadata={"source": "update", "step": 5},
|
||||
)
|
||||
|
||||
|
||||
@@ -2947,7 +2947,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value", "market": "DE"},
|
||||
next=("tool_two_slow",),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert await tool_two.ainvoke(None, thread1, debug=1) == {
|
||||
@@ -2958,7 +2958,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value slow", "market": "DE"},
|
||||
next=(),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
@@ -2971,7 +2971,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value", "market": "US"},
|
||||
next=("tool_two_fast",),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 0},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert await tool_two.ainvoke(None, thread2, debug=1) == {
|
||||
@@ -2982,7 +2982,7 @@ async def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
values={"my_key": "value fast", "market": "US"},
|
||||
next=(),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
|
||||
|
||||
@@ -3156,7 +3156,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared", "market": "DE"},
|
||||
next=("tool_two_slow",),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert await tool_two.ainvoke(None, thread1, debug=1) == {
|
||||
@@ -3167,7 +3167,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared slow finished", "market": "DE"},
|
||||
next=(),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
@@ -3180,7 +3180,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared", "market": "US"},
|
||||
next=("tool_two_fast",),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert await tool_two.ainvoke(None, thread2, debug=1) == {
|
||||
@@ -3191,7 +3191,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared fast finished", "market": "US"},
|
||||
next=(),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
|
||||
@@ -3213,7 +3213,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared", "market": "DE"},
|
||||
next=("tool_two_slow",),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert await tool_two.ainvoke(None, thread1, debug=1) == {
|
||||
@@ -3224,7 +3224,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared slow finished", "market": "DE"},
|
||||
next=(),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
@@ -3237,7 +3237,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared", "market": "US"},
|
||||
next=("tool_two_fast",),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 1},
|
||||
)
|
||||
# resume, for same result as above
|
||||
assert await tool_two.ainvoke(None, thread2, debug=1) == {
|
||||
@@ -3248,7 +3248,7 @@ async def test_branch_then() -> None:
|
||||
values={"my_key": "value prepared fast finished", "market": "US"},
|
||||
next=(),
|
||||
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
|
||||
metadata={},
|
||||
metadata={"source": "loop", "step": 3},
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user