mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-27 20:15:00 +02:00
Add metadata to checkpoints
- not yet used in this PR
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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"""
|
||||
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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={},
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user