Add checkpoint metadata fields

- source: input, update or loop
- step: int
- make step counter continue from previous last step
This commit is contained in:
Nuno Campos
2024-05-06 11:52:49 -07:00
parent 3ff3def62b
commit f160f82912
8 changed files with 160 additions and 109 deletions
+7 -6
View File
@@ -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()
+20 -3
View File
@@ -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
+9 -5
View File
@@ -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
)
+6 -5
View File
@@ -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 {
+42 -20
View File
@@ -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(
+8 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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},
)