Add metadata to checkpoints

- not yet used in this PR
This commit is contained in:
Nuno Campos
2024-05-06 11:52:49 -07:00
parent d79b483a67
commit f304908102
8 changed files with 188 additions and 84 deletions
+33 -36
View File
@@ -1,7 +1,7 @@
import asyncio
from contextlib import AbstractAsyncContextManager
from types import TracebackType
from typing import AsyncIterator, Optional
from typing import Any, AsyncIterator, Optional
import aiosqlite
from langchain_core.runnables import RunnableConfig
@@ -130,6 +130,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
thread_ts TEXT NOT NULL,
parent_ts TEXT,
checkpoint BLOB,
metadata BLOB,
PRIMARY KEY (thread_id, thread_ts)
);
"""
@@ -155,7 +156,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
await self.setup()
if config["configurable"].get("thread_ts"):
async with self.conn.execute(
"SELECT checkpoint, parent_ts FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
"SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
(
str(config["configurable"]["thread_id"]),
str(config["configurable"]["thread_ts"]),
@@ -165,20 +166,19 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
return CheckpointTuple(
config,
self.serde.loads(value[0]),
(
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"thread_ts": value[1],
}
self.serde.loads(value[2]) if value[2] is not None else None,
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"thread_ts": value[1],
}
if value[1]
else None
),
}
if value[1]
else None,
)
else:
async with self.conn.execute(
"SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1",
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1",
(str(config["configurable"]["thread_id"]),),
) as cursor:
if value := await cursor.fetchone():
@@ -190,16 +190,15 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
}
},
self.serde.loads(value[3]),
(
{
"configurable": {
"thread_id": value[0],
"thread_ts": value[2],
}
self.serde.loads(value[4]) if value[4] is not None else None,
{
"configurable": {
"thread_id": value[0],
"thread_ts": value[2],
}
if value[2]
else None
),
}
if value[2]
else None,
)
async def alist(
@@ -224,9 +223,9 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
"""
await self.setup()
query = (
"SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC"
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC"
if before is None
else "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC"
else "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC"
)
if limit:
query += f" LIMIT {limit}"
@@ -241,24 +240,21 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
)
),
) as cursor:
async for thread_id, thread_ts, parent_ts, value in cursor:
async for thread_id, thread_ts, parent_ts, value, metadata in cursor:
yield CheckpointTuple(
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
self.serde.loads(value),
(
{
"configurable": {
"thread_id": thread_id,
"thread_ts": parent_ts,
}
}
if parent_ts
else None
),
self.serde.loads(metadata) if metadata is not None else None,
{"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}}
if parent_ts
else None,
)
async def aput(
self, config: RunnableConfig, checkpoint: Checkpoint
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: Optional[dict[str, Any]] = None,
) -> RunnableConfig:
"""Save a checkpoint to the database asynchronously.
@@ -274,12 +270,13 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
"""
await self.setup()
async with self.conn.execute(
"INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint) VALUES (?, ?, ?, ?)",
"INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint, metadata) VALUES (?, ?, ?, ?, ?)",
(
str(config["configurable"]["thread_id"]),
checkpoint["ts"],
config["configurable"].get("thread_ts"),
self.serde.dumps(checkpoint),
self.serde.dumps(metadata) if metadata is not None else None,
),
):
await self.conn.commit()
+11 -2
View File
@@ -83,6 +83,7 @@ class CheckpointAt(StrEnum):
class CheckpointTuple(NamedTuple):
config: RunnableConfig
checkpoint: Checkpoint
metadata: Optional[dict[str, Any]]
parent_config: Optional[RunnableConfig] = None
@@ -139,7 +140,12 @@ class BaseCheckpointSaver(ABC):
) -> Iterator[CheckpointTuple]:
raise NotImplementedError
def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig:
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: dict[str, Any],
) -> RunnableConfig:
raise NotImplementedError
async def aget(self, config: RunnableConfig) -> Optional[Checkpoint]:
@@ -159,6 +165,9 @@ class BaseCheckpointSaver(ABC):
raise NotImplementedError
async def aput(
self, config: RunnableConfig, checkpoint: Checkpoint
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: dict[str, Any],
) -> RunnableConfig:
raise NotImplementedError
+24 -8
View File
@@ -1,6 +1,6 @@
import asyncio
from collections import defaultdict
from typing import AsyncIterator, Iterator, Optional
from typing import Any, AsyncIterator, Iterator, Optional
from langchain_core.runnables import RunnableConfig
@@ -39,7 +39,7 @@ class MemorySaver(BaseCheckpointSaver):
asyncio.run(coro) # Output: 2
"""
storage: defaultdict[str, dict[str, Checkpoint]]
storage: defaultdict[str, dict[str, tuple[bytes, bytes]]]
def __init__(
self,
@@ -66,16 +66,21 @@ class MemorySaver(BaseCheckpointSaver):
"""
thread_id = config["configurable"]["thread_id"]
if ts := config["configurable"].get("thread_ts"):
if checkpoint := self.storage[thread_id].get(ts):
if saved := self.storage[thread_id].get(ts):
checkpoint, metadata = saved
return CheckpointTuple(
config=config, checkpoint=self.serde.loads(checkpoint)
config=config,
checkpoint=self.serde.loads(checkpoint),
metadata=self.serde.loads(metadata),
)
else:
if checkpoints := self.storage[thread_id]:
ts = max(checkpoints.keys())
checkpoint, metadata = checkpoints[ts]
return CheckpointTuple(
config={"configurable": {"thread_id": thread_id, "thread_ts": ts}},
checkpoint=self.serde.loads(checkpoints[ts]),
checkpoint=self.serde.loads(checkpoint),
metadata=self.serde.loads(metadata),
)
def list(
@@ -99,7 +104,7 @@ class MemorySaver(BaseCheckpointSaver):
Iterator[CheckpointTuple]: An iterator of checkpoint tuples.
"""
thread_id = config["configurable"]["thread_id"]
for ts, checkpoint in self.storage[thread_id].items():
for ts, (checkpoint, metadata) in self.storage[thread_id].items():
if before and ts >= before["configurable"]["thread_ts"]:
continue
if limit is not None and limit <= 0:
@@ -108,9 +113,15 @@ class MemorySaver(BaseCheckpointSaver):
yield CheckpointTuple(
config={"configurable": {"thread_id": thread_id, "thread_ts": ts}},
checkpoint=self.serde.loads(checkpoint),
metadata=self.serde.loads(metadata),
)
def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig:
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: dict[str, Any] = None,
) -> RunnableConfig:
"""Save a checkpoint to the in-memory storage.
This method saves a checkpoint to the in-memory storage. The checkpoint is associated
@@ -124,7 +135,12 @@ class MemorySaver(BaseCheckpointSaver):
RunnableConfig: The updated config containing the saved checkpoint's timestamp.
"""
self.storage[config["configurable"]["thread_id"]].update(
{checkpoint["ts"]: self.serde.dumps(checkpoint)}
{
checkpoint["ts"]: (
self.serde.dumps(checkpoint),
self.serde.dumps(metadata or {}),
)
}
)
return {
"configurable": {
+39 -34
View File
@@ -146,6 +146,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
thread_ts TEXT NOT NULL,
parent_ts TEXT,
checkpoint BLOB,
metadata BLOB,
PRIMARY KEY (thread_id, thread_ts)
);
"""
@@ -211,7 +212,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
with self.cursor(transaction=False) as cur:
if config["configurable"].get("thread_ts"):
cur.execute(
"SELECT checkpoint, parent_ts FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
"SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
(
str(config["configurable"]["thread_id"]),
str(config["configurable"]["thread_ts"]),
@@ -221,20 +222,19 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
return CheckpointTuple(
config,
self.serde.loads(value[0]),
(
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"thread_ts": value[1],
}
self.serde.loads(value[2]) if value[2] is not None else None,
{
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"thread_ts": value[1],
}
if value[1]
else None
),
}
if value[1]
else None,
)
else:
cur.execute(
"SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1",
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1",
(str(config["configurable"]["thread_id"]),),
)
if value := cur.fetchone():
@@ -246,16 +246,15 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
}
},
self.serde.loads(value[3]),
(
{
"configurable": {
"thread_id": value[0],
"thread_ts": value[2],
}
self.serde.loads(value[4]) if value[4] is not None else None,
{
"configurable": {
"thread_id": value[0],
"thread_ts": value[2],
}
if value[2]
else None
),
}
if value[2]
else None,
)
def list(
@@ -289,9 +288,9 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
print(checkpoints) # Output: [CheckpointTuple(...), ...]
"""
query = (
"SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC"
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC"
if before is None
else "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC"
else "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts < ? ORDER BY thread_ts DESC"
)
if limit:
query += f" LIMIT {limit}"
@@ -307,23 +306,27 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
)
),
)
for thread_id, thread_ts, parent_ts, value in cur:
for thread_id, thread_ts, parent_ts, value, metadata in cur:
yield CheckpointTuple(
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
self.serde.loads(value),
(
{
"configurable": {
"thread_id": thread_id,
"thread_ts": parent_ts,
}
self.serde.loads(metadata) if metadata is not None else None,
{
"configurable": {
"thread_id": thread_id,
"thread_ts": parent_ts,
}
if parent_ts
else None
),
}
if parent_ts
else None,
)
def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig:
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: dict[str, Any] = None,
) -> RunnableConfig:
"""Save a checkpoint to the database.
This method saves a checkpoint to the SQLite database. The checkpoint is associated
@@ -332,6 +335,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
Args:
config (RunnableConfig): The config to associate with the checkpoint.
checkpoint (Checkpoint): The checkpoint to save.
metadata (Optional[dict[str, Any]]): Additional metadata to save with the checkpoint. Defaults to None.
Returns:
RunnableConfig: The updated config containing the saved checkpoint's timestamp.
@@ -347,12 +351,13 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
"""
with self.cursor() as cur:
cur.execute(
"INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint) VALUES (?, ?, ?, ?)",
"INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint, metadata) VALUES (?, ?, ?, ?, ?)",
(
str(config["configurable"]["thread_id"]),
checkpoint["ts"],
config["configurable"].get("thread_ts"),
self.serde.dumps(checkpoint),
self.serde.dumps(metadata) if metadata is not None else None,
),
)
return {
+11 -4
View File
@@ -343,6 +343,7 @@ class Pregel(
read_channels(channels, self.stream_channels_asis),
tuple(name for name, _ in next_tasks),
config,
saved.metadata or {},
)
async def aget_state(self, config: RunnableConfig) -> StateSnapshot:
@@ -361,6 +362,7 @@ class Pregel(
read_channels(channels, self.stream_channels_asis),
tuple(name for name, _ in next_tasks),
config,
saved.metadata or {},
)
def get_state_history(
@@ -374,7 +376,7 @@ class Pregel(
if not self.checkpointer:
raise ValueError("No checkpointer set")
for config, checkpoint, parent_config in self.checkpointer.list(
for config, checkpoint, metadata, parent_config in self.checkpointer.list(
config, before=before, limit=limit
):
with ChannelsManager(self.channels, checkpoint) as channels:
@@ -385,6 +387,7 @@ class Pregel(
read_channels(channels, self.stream_channels_asis),
tuple(name for name, _ in next_tasks),
config,
metadata or {},
parent_config,
)
@@ -399,9 +402,12 @@ class Pregel(
if not self.checkpointer:
raise ValueError("No checkpointer set")
async for config, checkpoint, parent_config in self.checkpointer.alist(
config, before=before, limit=limit
):
async for (
config,
checkpoint,
metadata,
parent_config,
) in self.checkpointer.alist(config, before=before, limit=limit):
async with AsyncChannelsManager(self.channels, checkpoint) as channels:
_, next_tasks = _prepare_next_tasks(
checkpoint, self.nodes, channels, for_execution=False
@@ -410,6 +416,7 @@ class Pregel(
read_channels(channels, self.stream_channels_asis),
tuple(name for name, _ in next_tasks),
config,
metadata or {},
parent_config,
)
+2
View File
@@ -25,6 +25,8 @@ class StateSnapshot(NamedTuple):
"""Nodes to execute in the next step, if any"""
config: RunnableConfig
"""Config used to fetch this snapshot"""
metadata: dict[str, Any]
"""Metadata associated with this snapshot"""
parent_config: Optional[RunnableConfig] = None
"""Config used to fetch the parent snapshot, if any"""
+39
View File
@@ -1228,6 +1228,7 @@ def test_conditional_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert (
app_w_interrupt.checkpointer.get_tuple(config).config["configurable"][
@@ -1261,6 +1262,7 @@ def test_conditional_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1344,6 +1346,7 @@ def test_conditional_graph(
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# test state get/update methods with interrupt_before
@@ -1379,6 +1382,7 @@ def test_conditional_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
app_w_interrupt.update_state(
@@ -1406,6 +1410,7 @@ def test_conditional_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1489,6 +1494,7 @@ def test_conditional_graph(
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# test re-invoke to continue with interrupt_before
@@ -1524,6 +1530,7 @@ def test_conditional_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1860,6 +1867,7 @@ def test_conditional_state_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
app_w_interrupt.update_state(
@@ -1885,6 +1893,7 @@ def test_conditional_state_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1943,6 +1952,7 @@ def test_conditional_state_graph(
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# test state get/update methods with interrupt_before
@@ -1977,6 +1987,7 @@ def test_conditional_state_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
app_w_interrupt.update_state(
@@ -2002,6 +2013,7 @@ def test_conditional_state_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2060,6 +2072,7 @@ def test_conditional_state_graph(
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# test w interrupt before all
@@ -2082,6 +2095,7 @@ def test_conditional_state_graph(
},
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2104,6 +2118,7 @@ def test_conditional_state_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2142,6 +2157,7 @@ def test_conditional_state_graph(
},
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2186,6 +2202,7 @@ def test_conditional_state_graph(
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2224,6 +2241,7 @@ def test_conditional_state_graph(
},
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -3042,6 +3060,7 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# modify ai message
@@ -3067,6 +3086,7 @@ def test_message_graph(
],
next=("action",),
config=next_config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -3132,6 +3152,7 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
app_w_interrupt.update_state(
@@ -3167,6 +3188,7 @@ def test_message_graph(
],
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
app_w_interrupt = workflow.compile(
@@ -3212,6 +3234,7 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# modify ai message
@@ -3240,6 +3263,7 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -3305,6 +3329,7 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
app_w_interrupt.update_state(
@@ -3340,6 +3365,7 @@ def test_message_graph(
],
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
# add an extra message as if it came from "action" node
@@ -3375,6 +3401,7 @@ def test_message_graph(
],
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
@@ -3503,6 +3530,7 @@ def test_start_branch_then(
values={"my_key": "value", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={},
)
# resume, for same result as above
assert tool_two.invoke(None, thread1, debug=1) == {
@@ -3513,6 +3541,7 @@ def test_start_branch_then(
values={"my_key": "value slow", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={},
)
thread2 = {"configurable": {"thread_id": "2"}}
@@ -3525,6 +3554,7 @@ def test_start_branch_then(
values={"my_key": "value", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={},
)
# resume, for same result as above
assert tool_two.invoke(None, thread2, debug=1) == {
@@ -3535,6 +3565,7 @@ def test_start_branch_then(
values={"my_key": "value fast", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={},
)
@@ -3855,6 +3886,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={},
)
# resume, for same result as above
assert tool_two.invoke(None, thread1, debug=1) == {
@@ -3865,6 +3897,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={},
)
thread2 = {"configurable": {"thread_id": "2"}}
@@ -3877,6 +3910,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={},
)
# resume, for same result as above
assert tool_two.invoke(None, thread2, debug=1) == {
@@ -3887,6 +3921,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={},
)
with SqliteSaver.from_conn_string(":memory:") as saver:
@@ -3909,6 +3944,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={},
)
# resume, for same result as above
assert tool_two.invoke(None, thread1, debug=1) == {
@@ -3919,6 +3955,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={},
)
thread2 = {"configurable": {"thread_id": "2"}}
@@ -3931,6 +3968,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={},
)
# resume, for same result as above
assert tool_two.invoke(None, thread2, debug=1) == {
@@ -3941,6 +3979,7 @@ def test_branch_then(snapshot: SnapshotAssertion, checkpoint_at: CheckpointAt) -
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={},
)
+29
View File
@@ -1306,6 +1306,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
await app_w_interrupt.aupdate_state(
@@ -1333,6 +1334,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -1416,6 +1418,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
# test state get/update methods with interrupt_before
@@ -1454,6 +1457,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
await app_w_interrupt.aupdate_state(
@@ -1481,6 +1485,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -1564,6 +1569,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
# test re-invoke to continue with interrupt_before
@@ -1602,6 +1608,7 @@ async def test_conditional_graph(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -1924,6 +1931,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
await app_w_interrupt.aupdate_state(
@@ -1949,6 +1957,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -2007,6 +2016,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
# test state get/update methods with interrupt_before
@@ -2043,6 +2053,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
await app_w_interrupt.aupdate_state(
@@ -2068,6 +2079,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -2126,6 +2138,7 @@ async def test_conditional_graph_state(checkpoint_at: CheckpointAt) -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
@@ -2734,6 +2747,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None:
],
next=("action",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={},
)
# modify ai message
@@ -2761,6 +2775,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None:
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -2813,6 +2828,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None:
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
await app_w_interrupt.aupdate_state(
@@ -2846,6 +2862,7 @@ async def test_message_graph(checkpoint_at: CheckpointAt) -> None:
],
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={},
)
@@ -2968,6 +2985,7 @@ async def test_start_branch_then(
values={"my_key": "value", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={},
)
# resume, for same result as above
assert await tool_two.ainvoke(None, thread1, debug=1) == {
@@ -2978,6 +2996,7 @@ async def test_start_branch_then(
values={"my_key": "value slow", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={},
)
thread2 = {"configurable": {"thread_id": "2"}}
@@ -2990,6 +3009,7 @@ async def test_start_branch_then(
values={"my_key": "value", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={},
)
# resume, for same result as above
assert await tool_two.ainvoke(None, thread2, debug=1) == {
@@ -3000,6 +3020,7 @@ async def test_start_branch_then(
values={"my_key": "value fast", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={},
)
@@ -3308,6 +3329,7 @@ async def test_branch_then(
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={},
)
# resume, for same result as above
assert await tool_two.ainvoke(None, thread1, debug=1) == {
@@ -3318,6 +3340,7 @@ async def test_branch_then(
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={},
)
thread2 = {"configurable": {"thread_id": "2"}}
@@ -3330,6 +3353,7 @@ async def test_branch_then(
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={},
)
# resume, for same result as above
assert await tool_two.ainvoke(None, thread2, debug=1) == {
@@ -3340,6 +3364,7 @@ async def test_branch_then(
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={},
)
async with AsyncSqliteSaver.from_conn_string(":memory:") as saver:
@@ -3362,6 +3387,7 @@ async def test_branch_then(
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={},
)
# resume, for same result as above
assert await tool_two.ainvoke(None, thread1, debug=1) == {
@@ -3372,6 +3398,7 @@ async def test_branch_then(
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={},
)
thread2 = {"configurable": {"thread_id": "2"}}
@@ -3384,6 +3411,7 @@ async def test_branch_then(
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={},
)
# resume, for same result as above
assert await tool_two.ainvoke(None, thread2, debug=1) == {
@@ -3394,6 +3422,7 @@ async def test_branch_then(
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={},
)