mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-08 02:37:52 +02:00
lib: Checkpoint pending writes whenever a node finishes (#976)
* lib: Checkpoint pending writes whenever a node finishes - Whenever a node finishes, checkpoint pending writes * Rename arg * Add tests, resume from pending writes * Add comment * Add descriptive error * Fix bug found by will * Fix comments * Lint * Don't save pending write if executing only one node in step
This commit is contained in:
@@ -2,7 +2,16 @@ import asyncio
|
|||||||
import functools
|
import functools
|
||||||
from contextlib import AbstractAsyncContextManager
|
from contextlib import AbstractAsyncContextManager
|
||||||
from types import TracebackType
|
from types import TracebackType
|
||||||
from typing import Any, AsyncIterator, Dict, Iterator, Optional, TypeVar
|
from typing import (
|
||||||
|
Any,
|
||||||
|
AsyncIterator,
|
||||||
|
Dict,
|
||||||
|
Iterator,
|
||||||
|
Optional,
|
||||||
|
Sequence,
|
||||||
|
Tuple,
|
||||||
|
TypeVar,
|
||||||
|
)
|
||||||
|
|
||||||
import aiosqlite
|
import aiosqlite
|
||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
@@ -203,6 +212,15 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
|||||||
metadata BLOB,
|
metadata BLOB,
|
||||||
PRIMARY KEY (thread_id, thread_ts)
|
PRIMARY KEY (thread_id, thread_ts)
|
||||||
);
|
);
|
||||||
|
CREATE TABLE IF NOT EXISTS writes (
|
||||||
|
thread_id TEXT NOT NULL,
|
||||||
|
thread_ts TEXT NOT NULL,
|
||||||
|
task_id TEXT NOT NULL,
|
||||||
|
idx INTEGER NOT NULL,
|
||||||
|
channel TEXT NOT NULL,
|
||||||
|
value BLOB,
|
||||||
|
PRIMARY KEY (thread_id, thread_ts, task_id, idx)
|
||||||
|
);
|
||||||
"""
|
"""
|
||||||
):
|
):
|
||||||
await self.conn.commit()
|
await self.conn.commit()
|
||||||
@@ -224,56 +242,58 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
|||||||
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
|
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
|
||||||
"""
|
"""
|
||||||
await self.setup()
|
await self.setup()
|
||||||
if config["configurable"].get("thread_ts"):
|
async with self.conn.cursor() as cur:
|
||||||
async with self.conn.execute(
|
# find the latest checkpoint for the thread_id
|
||||||
"SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
|
if config["configurable"].get("thread_ts"):
|
||||||
(
|
await cur.execute(
|
||||||
str(config["configurable"]["thread_id"]),
|
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
|
||||||
str(config["configurable"]["thread_ts"]),
|
(
|
||||||
),
|
str(config["configurable"]["thread_id"]),
|
||||||
) as cursor:
|
str(config["configurable"]["thread_ts"]),
|
||||||
if value := await cursor.fetchone():
|
),
|
||||||
return CheckpointTuple(
|
)
|
||||||
config,
|
else:
|
||||||
self.serde.loads(value[0]),
|
await cur.execute(
|
||||||
self.serde.loads(value[2]) if value[2] is not None else {},
|
"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"]),),
|
||||||
{
|
)
|
||||||
"configurable": {
|
# if a checkpoint is found, return it
|
||||||
"thread_id": config["configurable"]["thread_id"],
|
if value := await cur.fetchone():
|
||||||
"thread_ts": value[1],
|
if not config["configurable"].get("thread_ts"):
|
||||||
}
|
config = {
|
||||||
}
|
"configurable": {
|
||||||
if value[1]
|
"thread_id": value[0],
|
||||||
else None
|
"thread_ts": value[1],
|
||||||
),
|
}
|
||||||
)
|
}
|
||||||
else:
|
# find any pending writes
|
||||||
async with self.conn.execute(
|
await cur.execute(
|
||||||
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1",
|
"SELECT task_id, channel, value FROM writes WHERE thread_id = ? AND thread_ts = ?",
|
||||||
(str(config["configurable"]["thread_id"]),),
|
(
|
||||||
) as cursor:
|
str(config["configurable"]["thread_id"]),
|
||||||
if value := await cursor.fetchone():
|
str(config["configurable"]["thread_ts"]),
|
||||||
return CheckpointTuple(
|
),
|
||||||
|
)
|
||||||
|
# deserialize the checkpoint and metadata
|
||||||
|
return CheckpointTuple(
|
||||||
|
config,
|
||||||
|
self.serde.loads(value[3]),
|
||||||
|
self.serde.loads(value[4]) if value[4] is not None else {},
|
||||||
|
(
|
||||||
{
|
{
|
||||||
"configurable": {
|
"configurable": {
|
||||||
"thread_id": value[0],
|
"thread_id": value[0],
|
||||||
"thread_ts": value[1],
|
"thread_ts": value[2],
|
||||||
}
|
}
|
||||||
},
|
}
|
||||||
self.serde.loads(value[3]),
|
if value[2]
|
||||||
self.serde.loads(value[4]) if value[4] is not None else {},
|
else None
|
||||||
(
|
),
|
||||||
{
|
[
|
||||||
"configurable": {
|
(task_id, channel, self.serde.loads(value))
|
||||||
"thread_id": value[0],
|
async for task_id, channel, value in cur
|
||||||
"thread_ts": value[2],
|
],
|
||||||
}
|
)
|
||||||
}
|
|
||||||
if value[2]
|
|
||||||
else None
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
async def alist(
|
async def alist(
|
||||||
self,
|
self,
|
||||||
@@ -358,3 +378,26 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
|||||||
"thread_ts": checkpoint["id"],
|
"thread_ts": checkpoint["id"],
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async def aput_writes(
|
||||||
|
self,
|
||||||
|
config: RunnableConfig,
|
||||||
|
writes: Sequence[Tuple[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
) -> None:
|
||||||
|
await self.setup()
|
||||||
|
async with self.conn.executemany(
|
||||||
|
"INSERT OR REPLACE INTO writes (thread_id, thread_ts, task_id, idx, channel, value) VALUES (?, ?, ?, ?, ?, ?)",
|
||||||
|
[
|
||||||
|
(
|
||||||
|
str(config["configurable"]["thread_id"]),
|
||||||
|
str(config["configurable"]["thread_ts"]),
|
||||||
|
task_id,
|
||||||
|
idx,
|
||||||
|
channel,
|
||||||
|
self.serde.dumps(value),
|
||||||
|
)
|
||||||
|
for idx, (channel, value) in enumerate(writes)
|
||||||
|
],
|
||||||
|
):
|
||||||
|
await self.conn.commit()
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from typing import (
|
|||||||
Literal,
|
Literal,
|
||||||
NamedTuple,
|
NamedTuple,
|
||||||
Optional,
|
Optional,
|
||||||
|
Tuple,
|
||||||
TypedDict,
|
TypedDict,
|
||||||
TypeVar,
|
TypeVar,
|
||||||
Union,
|
Union,
|
||||||
@@ -117,6 +118,7 @@ class CheckpointTuple(NamedTuple):
|
|||||||
checkpoint: Checkpoint
|
checkpoint: Checkpoint
|
||||||
metadata: CheckpointMetadata
|
metadata: CheckpointMetadata
|
||||||
parent_config: Optional[RunnableConfig] = None
|
parent_config: Optional[RunnableConfig] = None
|
||||||
|
pending_writes: Optional[List[Tuple[str, str, Any]]] = None
|
||||||
|
|
||||||
|
|
||||||
CheckpointThreadId = ConfigurableFieldSpec(
|
CheckpointThreadId = ConfigurableFieldSpec(
|
||||||
@@ -177,6 +179,16 @@ class BaseCheckpointSaver(ABC):
|
|||||||
) -> RunnableConfig:
|
) -> RunnableConfig:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def put_writes(
|
||||||
|
self,
|
||||||
|
config: RunnableConfig,
|
||||||
|
writes: List[Tuple[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
) -> None:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"This method was added in langgraph 0.1.7. Please update your checkpointer to implement it."
|
||||||
|
)
|
||||||
|
|
||||||
async def aget(self, config: RunnableConfig) -> Optional[Checkpoint]:
|
async def aget(self, config: RunnableConfig) -> Optional[Checkpoint]:
|
||||||
if value := await self.aget_tuple(config):
|
if value := await self.aget_tuple(config):
|
||||||
return value.checkpoint
|
return value.checkpoint
|
||||||
@@ -203,6 +215,16 @@ class BaseCheckpointSaver(ABC):
|
|||||||
) -> RunnableConfig:
|
) -> RunnableConfig:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
async def aput_writes(
|
||||||
|
self,
|
||||||
|
config: RunnableConfig,
|
||||||
|
writes: List[Tuple[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
) -> None:
|
||||||
|
raise NotImplementedError(
|
||||||
|
"This method was added in langgraph 0.1.7. Please update your checkpointer to implement it."
|
||||||
|
)
|
||||||
|
|
||||||
def get_next_version(self, current: Optional[V], channel: BaseChannel) -> V:
|
def get_next_version(self, current: Optional[V], channel: BaseChannel) -> V:
|
||||||
"""Get the next version of a channel. Default is to use int versions, incrementing by 1. If you override, you can use str/int/float versions,
|
"""Get the next version of a channel. Default is to use int versions, incrementing by 1. If you override, you can use str/int/float versions,
|
||||||
as long as they are monotonically increasing."""
|
as long as they are monotonically increasing."""
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from typing import Any, AsyncIterator, Dict, Iterator, Optional
|
from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Tuple
|
||||||
|
|
||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
|
|
||||||
@@ -53,6 +53,7 @@ class MemorySaver(BaseCheckpointSaver):
|
|||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(serde=serde)
|
super().__init__(serde=serde)
|
||||||
self.storage = defaultdict(dict)
|
self.storage = defaultdict(dict)
|
||||||
|
self.writes = defaultdict(list)
|
||||||
|
|
||||||
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||||
"""Get a checkpoint tuple from the in-memory storage.
|
"""Get a checkpoint tuple from the in-memory storage.
|
||||||
@@ -72,19 +73,27 @@ class MemorySaver(BaseCheckpointSaver):
|
|||||||
if ts := config["configurable"].get("thread_ts"):
|
if ts := config["configurable"].get("thread_ts"):
|
||||||
if saved := self.storage[thread_id].get(ts):
|
if saved := self.storage[thread_id].get(ts):
|
||||||
checkpoint, metadata = saved
|
checkpoint, metadata = saved
|
||||||
|
writes = self.writes[(thread_id, ts)]
|
||||||
return CheckpointTuple(
|
return CheckpointTuple(
|
||||||
config=config,
|
config=config,
|
||||||
checkpoint=self.serde.loads(checkpoint),
|
checkpoint=self.serde.loads(checkpoint),
|
||||||
metadata=self.serde.loads(metadata),
|
metadata=self.serde.loads(metadata),
|
||||||
|
pending_writes=[
|
||||||
|
(id, c, self.serde.loads(v)) for id, c, v in writes
|
||||||
|
],
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if checkpoints := self.storage[thread_id]:
|
if checkpoints := self.storage[thread_id]:
|
||||||
ts = max(checkpoints.keys())
|
ts = max(checkpoints.keys())
|
||||||
checkpoint, metadata = checkpoints[ts]
|
checkpoint, metadata = checkpoints[ts]
|
||||||
|
writes = self.writes[(thread_id, ts)]
|
||||||
return CheckpointTuple(
|
return CheckpointTuple(
|
||||||
config={"configurable": {"thread_id": thread_id, "thread_ts": ts}},
|
config={"configurable": {"thread_id": thread_id, "thread_ts": ts}},
|
||||||
checkpoint=self.serde.loads(checkpoint),
|
checkpoint=self.serde.loads(checkpoint),
|
||||||
metadata=self.serde.loads(metadata),
|
metadata=self.serde.loads(metadata),
|
||||||
|
pending_writes=[
|
||||||
|
(id, c, self.serde.loads(v)) for id, c, v in writes
|
||||||
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
def list(
|
def list(
|
||||||
@@ -168,6 +177,30 @@ class MemorySaver(BaseCheckpointSaver):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def put_writes(
|
||||||
|
self,
|
||||||
|
config: RunnableConfig,
|
||||||
|
writes: List[Tuple[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
) -> RunnableConfig:
|
||||||
|
"""Save a list of writes to the in-memory storage.
|
||||||
|
|
||||||
|
This method saves a list of writes to the in-memory storage. The writes are associated
|
||||||
|
with the provided config.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config (RunnableConfig): The config to associate with the writes.
|
||||||
|
writes (list[tuple[str, Any]]): The writes to save.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
RunnableConfig: The updated config containing the saved writes' timestamp.
|
||||||
|
"""
|
||||||
|
thread_id = config["configurable"]["thread_id"]
|
||||||
|
ts = config["configurable"]["thread_ts"]
|
||||||
|
self.writes[(thread_id, ts)].extend(
|
||||||
|
[(task_id, c, self.serde.dumps(v)) for c, v in writes]
|
||||||
|
)
|
||||||
|
|
||||||
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||||
"""Asynchronous version of get_tuple.
|
"""Asynchronous version of get_tuple.
|
||||||
|
|
||||||
@@ -224,3 +257,13 @@ class MemorySaver(BaseCheckpointSaver):
|
|||||||
return await asyncio.get_running_loop().run_in_executor(
|
return await asyncio.get_running_loop().run_in_executor(
|
||||||
None, self.put, config, checkpoint, metadata
|
None, self.put, config, checkpoint, metadata
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def aput_writes(
|
||||||
|
self,
|
||||||
|
config: RunnableConfig,
|
||||||
|
writes: List[Tuple[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
) -> RunnableConfig:
|
||||||
|
return await asyncio.get_running_loop().run_in_executor(
|
||||||
|
None, self.put_writes, config, writes, task_id
|
||||||
|
)
|
||||||
|
|||||||
@@ -171,6 +171,15 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
|||||||
metadata BLOB,
|
metadata BLOB,
|
||||||
PRIMARY KEY (thread_id, thread_ts)
|
PRIMARY KEY (thread_id, thread_ts)
|
||||||
);
|
);
|
||||||
|
CREATE TABLE IF NOT EXISTS writes (
|
||||||
|
thread_id TEXT NOT NULL,
|
||||||
|
thread_ts TEXT NOT NULL,
|
||||||
|
task_id TEXT NOT NULL,
|
||||||
|
idx INTEGER NOT NULL,
|
||||||
|
channel TEXT NOT NULL,
|
||||||
|
value BLOB,
|
||||||
|
PRIMARY KEY (thread_id, thread_ts, task_id, idx)
|
||||||
|
);
|
||||||
"""
|
"""
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -233,56 +242,57 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
|||||||
CheckpointTuple(...)
|
CheckpointTuple(...)
|
||||||
""" # noqa
|
""" # noqa
|
||||||
with self.cursor(transaction=False) as cur:
|
with self.cursor(transaction=False) as cur:
|
||||||
|
# find the latest checkpoint for the thread_id
|
||||||
if config["configurable"].get("thread_ts"):
|
if config["configurable"].get("thread_ts"):
|
||||||
cur.execute(
|
cur.execute(
|
||||||
"SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
|
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?",
|
||||||
(
|
(
|
||||||
str(config["configurable"]["thread_id"]),
|
str(config["configurable"]["thread_id"]),
|
||||||
str(config["configurable"]["thread_ts"]),
|
str(config["configurable"]["thread_ts"]),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
if value := cur.fetchone():
|
|
||||||
return CheckpointTuple(
|
|
||||||
config,
|
|
||||||
self.serde.loads(value[0]),
|
|
||||||
self.serde.loads(value[2]) if value[2] is not None else {},
|
|
||||||
(
|
|
||||||
{
|
|
||||||
"configurable": {
|
|
||||||
"thread_id": config["configurable"]["thread_id"],
|
|
||||||
"thread_ts": value[1],
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if value[1]
|
|
||||||
else None
|
|
||||||
),
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
cur.execute(
|
cur.execute(
|
||||||
"SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata 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"]),),
|
(str(config["configurable"]["thread_id"]),),
|
||||||
)
|
)
|
||||||
if value := cur.fetchone():
|
# if a checkpoint is found, return it
|
||||||
return CheckpointTuple(
|
if value := cur.fetchone():
|
||||||
|
if not config["configurable"].get("thread_ts"):
|
||||||
|
config = {
|
||||||
|
"configurable": {
|
||||||
|
"thread_id": value[0],
|
||||||
|
"thread_ts": value[1],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
# find any pending writes
|
||||||
|
cur.execute(
|
||||||
|
"SELECT task_id, channel, value FROM writes WHERE thread_id = ? AND thread_ts = ?",
|
||||||
|
(
|
||||||
|
str(config["configurable"]["thread_id"]),
|
||||||
|
str(config["configurable"]["thread_ts"]),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
# deserialize the checkpoint and metadata
|
||||||
|
return CheckpointTuple(
|
||||||
|
config,
|
||||||
|
self.serde.loads(value[3]),
|
||||||
|
self.serde.loads(value[4]) if value[4] is not None else {},
|
||||||
|
(
|
||||||
{
|
{
|
||||||
"configurable": {
|
"configurable": {
|
||||||
"thread_id": value[0],
|
"thread_id": value[0],
|
||||||
"thread_ts": value[1],
|
"thread_ts": value[2],
|
||||||
}
|
}
|
||||||
},
|
}
|
||||||
self.serde.loads(value[3]),
|
if value[2]
|
||||||
self.serde.loads(value[4]) if value[4] is not None else {},
|
else None
|
||||||
(
|
),
|
||||||
{
|
[
|
||||||
"configurable": {
|
(task_id, channel, self.serde.loads(value))
|
||||||
"thread_id": value[0],
|
for task_id, channel, value in cur
|
||||||
"thread_ts": value[2],
|
],
|
||||||
}
|
)
|
||||||
}
|
|
||||||
if value[2]
|
|
||||||
else None
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def list(
|
def list(
|
||||||
self,
|
self,
|
||||||
@@ -394,6 +404,28 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def put_writes(
|
||||||
|
self,
|
||||||
|
config: RunnableConfig,
|
||||||
|
writes: Sequence[Tuple[str, Any]],
|
||||||
|
task_id: str,
|
||||||
|
) -> None:
|
||||||
|
with self.lock, self.cursor() as cur:
|
||||||
|
cur.executemany(
|
||||||
|
"INSERT OR REPLACE INTO writes (thread_id, thread_ts, task_id, idx, channel, value) VALUES (?, ?, ?, ?, ?, ?)",
|
||||||
|
[
|
||||||
|
(
|
||||||
|
str(config["configurable"]["thread_id"]),
|
||||||
|
str(config["configurable"]["thread_ts"]),
|
||||||
|
task_id,
|
||||||
|
idx,
|
||||||
|
channel,
|
||||||
|
self.serde.dumps(value),
|
||||||
|
)
|
||||||
|
for idx, (channel, value) in enumerate(writes)
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
||||||
"""Get a checkpoint tuple from the database asynchronously.
|
"""Get a checkpoint tuple from the database asynchronously.
|
||||||
|
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
|
import json
|
||||||
import time
|
import time
|
||||||
from collections import defaultdict, deque
|
from collections import defaultdict, deque
|
||||||
from functools import partial
|
from functools import partial
|
||||||
@@ -23,6 +24,7 @@ from typing import (
|
|||||||
get_type_hints,
|
get_type_hints,
|
||||||
overload,
|
overload,
|
||||||
)
|
)
|
||||||
|
from uuid import UUID, uuid5
|
||||||
|
|
||||||
from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager
|
from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager
|
||||||
from langchain_core.globals import get_debug
|
from langchain_core.globals import get_debug
|
||||||
@@ -442,7 +444,7 @@ class Pregel(
|
|||||||
and signature(self.checkpointer.list).parameters.get("filter") is None
|
and signature(self.checkpointer.list).parameters.get("filter") is None
|
||||||
):
|
):
|
||||||
raise ValueError("Checkpointer does not support filtering")
|
raise ValueError("Checkpointer does not support filtering")
|
||||||
for config, checkpoint, metadata, parent_config in self.checkpointer.list(
|
for config, checkpoint, metadata, parent_config, _ in self.checkpointer.list(
|
||||||
config, before=before, limit=limit, filter=filter
|
config, before=before, limit=limit, filter=filter
|
||||||
):
|
):
|
||||||
with ChannelsManager(
|
with ChannelsManager(
|
||||||
@@ -489,6 +491,7 @@ class Pregel(
|
|||||||
checkpoint,
|
checkpoint,
|
||||||
metadata,
|
metadata,
|
||||||
parent_config,
|
parent_config,
|
||||||
|
_,
|
||||||
) in self.checkpointer.alist(config, before=before, limit=limit, filter=filter):
|
) in self.checkpointer.alist(config, before=before, limit=limit, filter=filter):
|
||||||
async with AsyncChannelsManager(
|
async with AsyncChannelsManager(
|
||||||
self.channels, checkpoint, config
|
self.channels, checkpoint, config
|
||||||
@@ -565,6 +568,7 @@ class Pregel(
|
|||||||
deque(),
|
deque(),
|
||||||
None,
|
None,
|
||||||
[INTERRUPT],
|
[INTERRUPT],
|
||||||
|
str(uuid5(UUID(checkpoint["id"]), INTERRUPT)),
|
||||||
)
|
)
|
||||||
# execute task
|
# execute task
|
||||||
task.proc.invoke(
|
task.proc.invoke(
|
||||||
@@ -653,6 +657,7 @@ class Pregel(
|
|||||||
deque(),
|
deque(),
|
||||||
None,
|
None,
|
||||||
[INTERRUPT],
|
[INTERRUPT],
|
||||||
|
str(uuid5(UUID(checkpoint["id"]), INTERRUPT)),
|
||||||
)
|
)
|
||||||
# execute task
|
# execute task
|
||||||
await task.proc.ainvoke(
|
await task.proc.ainvoke(
|
||||||
@@ -879,6 +884,23 @@ class Pregel(
|
|||||||
self.managed_values_dict, config, self
|
self.managed_values_dict, config, self
|
||||||
) as managed:
|
) as managed:
|
||||||
|
|
||||||
|
def put_writes(task_id: str, writes: Sequence[tuple[str, Any]]) -> None:
|
||||||
|
if self.checkpointer is not None:
|
||||||
|
bg.append(
|
||||||
|
executor.submit(
|
||||||
|
self.checkpointer.put_writes,
|
||||||
|
{
|
||||||
|
**checkpoint_config,
|
||||||
|
"configurable": {
|
||||||
|
**checkpoint_config["configurable"],
|
||||||
|
"thread_ts": checkpoint["id"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
writes,
|
||||||
|
task_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def put_checkpoint(metadata: CheckpointMetadata) -> Iterator[Any]:
|
def put_checkpoint(metadata: CheckpointMetadata) -> Iterator[Any]:
|
||||||
nonlocal checkpoint, checkpoint_config, channels
|
nonlocal checkpoint, checkpoint_config, channels
|
||||||
|
|
||||||
@@ -963,8 +985,7 @@ class Pregel(
|
|||||||
# increment start to 0
|
# increment start to 0
|
||||||
start += 1
|
start += 1
|
||||||
else:
|
else:
|
||||||
# if received no input, take that as signal to proceed
|
# no input is taken as signal to proceed past previous interrupt
|
||||||
# past previous interrupt, if any
|
|
||||||
checkpoint = copy_checkpoint(checkpoint)
|
checkpoint = copy_checkpoint(checkpoint)
|
||||||
for k in self.stream_channels_list:
|
for k in self.stream_channels_list:
|
||||||
if k in checkpoint["channel_versions"]:
|
if k in checkpoint["channel_versions"]:
|
||||||
@@ -994,6 +1015,15 @@ class Pregel(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# assign pending writes to tasks
|
||||||
|
if saved and saved.pending_writes:
|
||||||
|
for task in next_tasks:
|
||||||
|
task.writes.extend(
|
||||||
|
(c, v)
|
||||||
|
for tid, c, v in saved.pending_writes
|
||||||
|
if tid == task.id
|
||||||
|
)
|
||||||
|
|
||||||
# if no more tasks, we're done
|
# if no more tasks, we're done
|
||||||
if not next_tasks:
|
if not next_tasks:
|
||||||
if step == start:
|
if step == start:
|
||||||
@@ -1027,12 +1057,15 @@ class Pregel(
|
|||||||
futures = {
|
futures = {
|
||||||
executor.submit(run_with_retry, task, self.retry_policy): task
|
executor.submit(run_with_retry, task, self.retry_policy): task
|
||||||
for task in next_tasks
|
for task in next_tasks
|
||||||
|
if not task.writes
|
||||||
}
|
}
|
||||||
end_time = (
|
end_time = (
|
||||||
self.step_timeout + time.monotonic()
|
self.step_timeout + time.monotonic()
|
||||||
if self.step_timeout
|
if self.step_timeout
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
if not futures:
|
||||||
|
done, inflight = set(), set()
|
||||||
while futures:
|
while futures:
|
||||||
done, inflight = concurrent.futures.wait(
|
done, inflight = concurrent.futures.wait(
|
||||||
futures,
|
futures,
|
||||||
@@ -1050,6 +1083,10 @@ class Pregel(
|
|||||||
# exception will be handled in panic_or_proceed
|
# exception will be handled in panic_or_proceed
|
||||||
futures.clear()
|
futures.clear()
|
||||||
else:
|
else:
|
||||||
|
# save task writes to checkpointer, unless this
|
||||||
|
# is the single or last task in this step
|
||||||
|
if futures:
|
||||||
|
put_writes(task.id, task.writes)
|
||||||
# yield updates output for the finished task
|
# yield updates output for the finished task
|
||||||
if "updates" in stream_modes:
|
if "updates" in stream_modes:
|
||||||
yield from _with_mode(
|
yield from _with_mode(
|
||||||
@@ -1076,7 +1113,7 @@ class Pregel(
|
|||||||
|
|
||||||
# combine pending writes from all tasks
|
# combine pending writes from all tasks
|
||||||
pending_writes = deque[tuple[str, Any]]()
|
pending_writes = deque[tuple[str, Any]]()
|
||||||
for _, _, _, writes, _, _ in next_tasks:
|
for _, _, _, writes, _, _, _ in next_tasks:
|
||||||
pending_writes.extend(writes)
|
pending_writes.extend(writes)
|
||||||
|
|
||||||
if debug:
|
if debug:
|
||||||
@@ -1240,6 +1277,24 @@ class Pregel(
|
|||||||
self.managed_values_dict, config, self
|
self.managed_values_dict, config, self
|
||||||
) as managed:
|
) as managed:
|
||||||
|
|
||||||
|
def put_writes(task_id: str, writes: Sequence[tuple[str, Any]]) -> None:
|
||||||
|
if self.checkpointer is not None:
|
||||||
|
bg.append(
|
||||||
|
asyncio.create_task(
|
||||||
|
self.checkpointer.aput_writes(
|
||||||
|
{
|
||||||
|
**checkpoint_config,
|
||||||
|
"configurable": {
|
||||||
|
**checkpoint_config["configurable"],
|
||||||
|
"thread_ts": checkpoint["id"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
writes,
|
||||||
|
task_id,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def put_checkpoint(metadata: CheckpointMetadata) -> Iterator[Any]:
|
def put_checkpoint(metadata: CheckpointMetadata) -> Iterator[Any]:
|
||||||
nonlocal checkpoint, checkpoint_config, channels
|
nonlocal checkpoint, checkpoint_config, channels
|
||||||
|
|
||||||
@@ -1320,8 +1375,7 @@ class Pregel(
|
|||||||
# increment start to 0
|
# increment start to 0
|
||||||
start += 1
|
start += 1
|
||||||
else:
|
else:
|
||||||
# if received no input, take that as signal to proceed
|
# no input is taken as signal to proceed past previous interrupt
|
||||||
# past previous interrupt, if any
|
|
||||||
checkpoint = copy_checkpoint(checkpoint)
|
checkpoint = copy_checkpoint(checkpoint)
|
||||||
for k in self.stream_channels_list:
|
for k in self.stream_channels_list:
|
||||||
if k in checkpoint["channel_versions"]:
|
if k in checkpoint["channel_versions"]:
|
||||||
@@ -1351,6 +1405,15 @@ class Pregel(
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# assign pending writes to tasks
|
||||||
|
if saved and saved.pending_writes:
|
||||||
|
for task in next_tasks:
|
||||||
|
task.writes.extend(
|
||||||
|
(c, v)
|
||||||
|
for tid, c, v in saved.pending_writes
|
||||||
|
if tid == task.id
|
||||||
|
)
|
||||||
|
|
||||||
# if no more tasks, we're done
|
# if no more tasks, we're done
|
||||||
if not next_tasks:
|
if not next_tasks:
|
||||||
if step == start:
|
if step == start:
|
||||||
@@ -1387,10 +1450,13 @@ class Pregel(
|
|||||||
arun_with_retry(task, self.retry_policy, do_stream)
|
arun_with_retry(task, self.retry_policy, do_stream)
|
||||||
): task
|
): task
|
||||||
for task in next_tasks
|
for task in next_tasks
|
||||||
|
if not task.writes
|
||||||
}
|
}
|
||||||
end_time = (
|
end_time = (
|
||||||
self.step_timeout + loop.time() if self.step_timeout else None
|
self.step_timeout + loop.time() if self.step_timeout else None
|
||||||
)
|
)
|
||||||
|
if not futures:
|
||||||
|
done, inflight = set(), set()
|
||||||
while futures:
|
while futures:
|
||||||
done, inflight = await asyncio.wait(
|
done, inflight = await asyncio.wait(
|
||||||
futures,
|
futures,
|
||||||
@@ -1406,6 +1472,10 @@ class Pregel(
|
|||||||
# exception will be handle in panic_or_proceed
|
# exception will be handle in panic_or_proceed
|
||||||
futures.clear()
|
futures.clear()
|
||||||
else:
|
else:
|
||||||
|
# save task writes to checkpointer, unless this
|
||||||
|
# is the single or last task in this step
|
||||||
|
if futures:
|
||||||
|
put_writes(task.id, task.writes)
|
||||||
# yield updates output for the finished task
|
# yield updates output for the finished task
|
||||||
if "updates" in stream_modes:
|
if "updates" in stream_modes:
|
||||||
for chunk in _with_mode(
|
for chunk in _with_mode(
|
||||||
@@ -1434,7 +1504,7 @@ class Pregel(
|
|||||||
|
|
||||||
# combine pending writes from all tasks
|
# combine pending writes from all tasks
|
||||||
pending_writes = deque[tuple[str, Any]]()
|
pending_writes = deque[tuple[str, Any]]()
|
||||||
for _, _, _, writes, _, _ in next_tasks:
|
for _, _, _, writes, _, _, _ in next_tasks:
|
||||||
pending_writes.extend(writes)
|
pending_writes.extend(writes)
|
||||||
|
|
||||||
if debug:
|
if debug:
|
||||||
@@ -1671,7 +1741,7 @@ def _should_interrupt(
|
|||||||
# and any triggered node is in interrupt_nodes list
|
# and any triggered node is in interrupt_nodes list
|
||||||
and any(
|
and any(
|
||||||
node
|
node
|
||||||
for node, _, _, _, config, _ in tasks
|
for node, _, _, _, config, _, _ in tasks
|
||||||
if (
|
if (
|
||||||
(not config or TAG_HIDDEN not in config.get("tags"))
|
(not config or TAG_HIDDEN not in config.get("tags"))
|
||||||
if interrupt_nodes == "*"
|
if interrupt_nodes == "*"
|
||||||
@@ -1825,6 +1895,14 @@ def _prepare_next_tasks(
|
|||||||
continue
|
continue
|
||||||
if for_execution:
|
if for_execution:
|
||||||
if node := processes[packet.node].get_node():
|
if node := processes[packet.node].get_node():
|
||||||
|
triggers = [TASKS]
|
||||||
|
metadata = {
|
||||||
|
"langgraph_step": step,
|
||||||
|
"langgraph_node": packet.node,
|
||||||
|
"langgraph_triggers": triggers,
|
||||||
|
"langgraph_task_idx": len(tasks),
|
||||||
|
}
|
||||||
|
task_id = str(uuid5(UUID(checkpoint["id"]), json.dumps(metadata)))
|
||||||
writes = deque()
|
writes = deque()
|
||||||
tasks.append(
|
tasks.append(
|
||||||
PregelExecutableTask(
|
PregelExecutableTask(
|
||||||
@@ -1836,14 +1914,7 @@ def _prepare_next_tasks(
|
|||||||
merge_configs(
|
merge_configs(
|
||||||
config,
|
config,
|
||||||
processes[packet.node].config,
|
processes[packet.node].config,
|
||||||
{
|
{"metadata": metadata},
|
||||||
"metadata": {
|
|
||||||
"langgraph_step": step,
|
|
||||||
"langgraph_node": packet.node,
|
|
||||||
"langgraph_triggers": [TASKS],
|
|
||||||
"langgraph_task_idx": len(tasks),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
),
|
),
|
||||||
run_name=packet.node,
|
run_name=packet.node,
|
||||||
callbacks=(
|
callbacks=(
|
||||||
@@ -1857,11 +1928,12 @@ def _prepare_next_tasks(
|
|||||||
_local_write, writes.extend, processes, channels
|
_local_write, writes.extend, processes, channels
|
||||||
),
|
),
|
||||||
CONFIG_KEY_READ: partial(
|
CONFIG_KEY_READ: partial(
|
||||||
_local_read, checkpoint, channels, tasks, config
|
_local_read, checkpoint, channels, writes, config
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
[TASKS],
|
triggers,
|
||||||
|
task_id,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -1879,7 +1951,7 @@ def _prepare_next_tasks(
|
|||||||
for name, proc in processes.items():
|
for name, proc in processes.items():
|
||||||
seen = checkpoint["versions_seen"][name]
|
seen = checkpoint["versions_seen"][name]
|
||||||
# If any of the channels read by this process were updated
|
# If any of the channels read by this process were updated
|
||||||
if triggers := [
|
if triggers := sorted(
|
||||||
chan
|
chan
|
||||||
for chan in proc.triggers
|
for chan in proc.triggers
|
||||||
if not isinstance(
|
if not isinstance(
|
||||||
@@ -1887,7 +1959,7 @@ def _prepare_next_tasks(
|
|||||||
)
|
)
|
||||||
and checkpoint["channel_versions"].get(chan, null_version)
|
and checkpoint["channel_versions"].get(chan, null_version)
|
||||||
> seen.get(chan, null_version)
|
> seen.get(chan, null_version)
|
||||||
]:
|
):
|
||||||
channels_to_consume.update(triggers)
|
channels_to_consume.update(triggers)
|
||||||
try:
|
try:
|
||||||
val = next(_proc_input(step, name, proc, managed, channels))
|
val = next(_proc_input(step, name, proc, managed, channels))
|
||||||
@@ -1906,8 +1978,14 @@ def _prepare_next_tasks(
|
|||||||
|
|
||||||
if for_execution:
|
if for_execution:
|
||||||
if node := proc.get_node():
|
if node := proc.get_node():
|
||||||
|
metadata = {
|
||||||
|
"langgraph_step": step,
|
||||||
|
"langgraph_node": name,
|
||||||
|
"langgraph_triggers": triggers,
|
||||||
|
"langgraph_task_idx": len(tasks),
|
||||||
|
}
|
||||||
|
task_id = str(uuid5(UUID(checkpoint["id"]), json.dumps(metadata)))
|
||||||
writes = deque()
|
writes = deque()
|
||||||
triggers = sorted(triggers)
|
|
||||||
tasks.append(
|
tasks.append(
|
||||||
PregelExecutableTask(
|
PregelExecutableTask(
|
||||||
name,
|
name,
|
||||||
@@ -1918,14 +1996,7 @@ def _prepare_next_tasks(
|
|||||||
merge_configs(
|
merge_configs(
|
||||||
config,
|
config,
|
||||||
proc.config,
|
proc.config,
|
||||||
{
|
{"metadata": metadata},
|
||||||
"metadata": {
|
|
||||||
"langgraph_step": step,
|
|
||||||
"langgraph_node": name,
|
|
||||||
"langgraph_triggers": triggers,
|
|
||||||
"langgraph_task_idx": len(tasks),
|
|
||||||
}
|
|
||||||
},
|
|
||||||
),
|
),
|
||||||
run_name=name,
|
run_name=name,
|
||||||
callbacks=(
|
callbacks=(
|
||||||
@@ -1948,6 +2019,7 @@ def _prepare_next_tasks(
|
|||||||
},
|
},
|
||||||
),
|
),
|
||||||
triggers,
|
triggers,
|
||||||
|
task_id,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ def map_debug_tasks(
|
|||||||
step: int, tasks: list[PregelExecutableTask]
|
step: int, tasks: list[PregelExecutableTask]
|
||||||
) -> Iterator[DebugOutputTask]:
|
) -> Iterator[DebugOutputTask]:
|
||||||
ts = datetime.now(timezone.utc).isoformat()
|
ts = datetime.now(timezone.utc).isoformat()
|
||||||
for name, input, _, _, config, triggers in tasks:
|
for name, input, _, _, config, triggers, _ in tasks:
|
||||||
if config is not None and TAG_HIDDEN in config.get("tags", []):
|
if config is not None and TAG_HIDDEN in config.get("tags", []):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -91,7 +91,7 @@ def map_debug_task_results(
|
|||||||
stream_channels_list: Sequence[str],
|
stream_channels_list: Sequence[str],
|
||||||
) -> Iterator[DebugOutputTaskResult]:
|
) -> Iterator[DebugOutputTaskResult]:
|
||||||
ts = datetime.now(timezone.utc).isoformat()
|
ts = datetime.now(timezone.utc).isoformat()
|
||||||
for name, _, _, writes, config, _ in tasks:
|
for name, _, _, writes, config, _, _ in tasks:
|
||||||
if config is not None and TAG_HIDDEN in config.get("tags", []):
|
if config is not None and TAG_HIDDEN in config.get("tags", []):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -138,7 +138,7 @@ def print_step_tasks(step: int, next_tasks: list[PregelExecutableTask]) -> None:
|
|||||||
)
|
)
|
||||||
+ "\n".join(
|
+ "\n".join(
|
||||||
f"- {get_colored_text(name, 'green')} -> {pformat(val)}"
|
f"- {get_colored_text(name, 'green')} -> {pformat(val)}"
|
||||||
for name, val, _, _, _, _ in next_tasks
|
for name, val, _, _, _, _, _ in next_tasks
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -105,7 +105,7 @@ def map_output_updates(
|
|||||||
if isinstance(output_channels, str):
|
if isinstance(output_channels, str):
|
||||||
if updated := [
|
if updated := [
|
||||||
(node, value)
|
(node, value)
|
||||||
for node, _, _, writes, _, _ in output_tasks
|
for node, _, _, writes, _, _, _ in output_tasks
|
||||||
for chan, value in writes
|
for chan, value in writes
|
||||||
if chan == output_channels
|
if chan == output_channels
|
||||||
]:
|
]:
|
||||||
@@ -122,7 +122,7 @@ def map_output_updates(
|
|||||||
node,
|
node,
|
||||||
{chan: value for chan, value in writes if chan in output_channels},
|
{chan: value for chan, value in writes if chan in output_channels},
|
||||||
)
|
)
|
||||||
for node, _, _, writes, _, _ in output_tasks
|
for node, _, _, writes, _, _, _ in output_tasks
|
||||||
if any(chan in output_channels for chan, _ in writes)
|
if any(chan in output_channels for chan, _ in writes)
|
||||||
]:
|
]:
|
||||||
grouped = defaultdict(list)
|
grouped = defaultdict(list)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ class PregelExecutableTask(NamedTuple):
|
|||||||
writes: deque[tuple[str, Any]]
|
writes: deque[tuple[str, Any]]
|
||||||
config: RunnableConfig
|
config: RunnableConfig
|
||||||
triggers: list[str]
|
triggers: list[str]
|
||||||
|
id: str
|
||||||
|
|
||||||
|
|
||||||
class StateSnapshot(NamedTuple):
|
class StateSnapshot(NamedTuple):
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from typing import (
|
|||||||
Literal,
|
Literal,
|
||||||
Optional,
|
Optional,
|
||||||
Sequence,
|
Sequence,
|
||||||
|
Tuple,
|
||||||
TypedDict,
|
TypedDict,
|
||||||
Union,
|
Union,
|
||||||
)
|
)
|
||||||
@@ -37,6 +38,7 @@ from langgraph.channels.context import Context
|
|||||||
from langgraph.channels.last_value import LastValue
|
from langgraph.channels.last_value import LastValue
|
||||||
from langgraph.channels.topic import Topic
|
from langgraph.channels.topic import Topic
|
||||||
from langgraph.checkpoint.base import (
|
from langgraph.checkpoint.base import (
|
||||||
|
BaseCheckpointSaver,
|
||||||
Checkpoint,
|
Checkpoint,
|
||||||
CheckpointMetadata,
|
CheckpointMetadata,
|
||||||
CheckpointTuple,
|
CheckpointTuple,
|
||||||
@@ -193,6 +195,12 @@ def test_checkpoint_errors() -> None:
|
|||||||
) -> RunnableConfig:
|
) -> RunnableConfig:
|
||||||
raise ValueError("Faulty put")
|
raise ValueError("Faulty put")
|
||||||
|
|
||||||
|
class FaultyPutWritesCheckpointer(MemorySaver):
|
||||||
|
def put_writes(
|
||||||
|
self, config: RunnableConfig, writes: List[Tuple[str, Any]], task_id: str
|
||||||
|
) -> RunnableConfig:
|
||||||
|
raise ValueError("Faulty put_writes")
|
||||||
|
|
||||||
class FaultyVersionCheckpointer(MemorySaver):
|
class FaultyVersionCheckpointer(MemorySaver):
|
||||||
def get_next_version(self, current: Optional[int], channel: BaseChannel) -> int:
|
def get_next_version(self, current: Optional[int], channel: BaseChannel) -> int:
|
||||||
raise ValueError("Faulty get_next_version")
|
raise ValueError("Faulty get_next_version")
|
||||||
@@ -200,10 +208,9 @@ def test_checkpoint_errors() -> None:
|
|||||||
def logic(inp: str) -> str:
|
def logic(inp: str) -> str:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
builder = Graph()
|
builder = StateGraph(Annotated[str, operator.add])
|
||||||
builder.add_node("agent", logic)
|
builder.add_node("agent", logic)
|
||||||
builder.set_entry_point("agent")
|
builder.add_edge(START, "agent")
|
||||||
builder.set_finish_point("agent")
|
|
||||||
|
|
||||||
graph = builder.compile(checkpointer=FaultyGetCheckpointer())
|
graph = builder.compile(checkpointer=FaultyGetCheckpointer())
|
||||||
with pytest.raises(ValueError, match="Faulty get_tuple"):
|
with pytest.raises(ValueError, match="Faulty get_tuple"):
|
||||||
@@ -217,6 +224,13 @@ def test_checkpoint_errors() -> None:
|
|||||||
with pytest.raises(ValueError, match="Faulty get_next_version"):
|
with pytest.raises(ValueError, match="Faulty get_next_version"):
|
||||||
graph.invoke("", {"configurable": {"thread_id": "thread-1"}})
|
graph.invoke("", {"configurable": {"thread_id": "thread-1"}})
|
||||||
|
|
||||||
|
# add parallel node
|
||||||
|
builder.add_node("parallel", logic)
|
||||||
|
builder.add_edge(START, "parallel")
|
||||||
|
graph = builder.compile(checkpointer=FaultyPutWritesCheckpointer())
|
||||||
|
with pytest.raises(ValueError, match="Faulty put_writes"):
|
||||||
|
graph.invoke("", {"configurable": {"thread_id": "thread-1"}})
|
||||||
|
|
||||||
|
|
||||||
def test_reducer_before_first_node() -> None:
|
def test_reducer_before_first_node() -> None:
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
@@ -944,6 +958,122 @@ def test_invoke_checkpoint(mocker: MockerFixture) -> None:
|
|||||||
assert checkpoint["channel_values"].get("total") == 5
|
assert checkpoint["channel_values"].get("total") == 5
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"checkpointer",
|
||||||
|
[
|
||||||
|
MemorySaverAssertImmutable(),
|
||||||
|
SqliteSaver.from_conn_string(":memory:"),
|
||||||
|
],
|
||||||
|
ids=[
|
||||||
|
"memory",
|
||||||
|
"sqlite",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None:
|
||||||
|
try:
|
||||||
|
|
||||||
|
class State(TypedDict):
|
||||||
|
value: Annotated[int, operator.add]
|
||||||
|
|
||||||
|
class AwhileMaker:
|
||||||
|
def __init__(self, sleep: float, rtn: Union[Dict, Exception]) -> None:
|
||||||
|
self.sleep = sleep
|
||||||
|
self.rtn = rtn
|
||||||
|
self.reset()
|
||||||
|
|
||||||
|
def __call__(self, input: State) -> Any:
|
||||||
|
self.calls += 1
|
||||||
|
time.sleep(self.sleep)
|
||||||
|
if isinstance(self.rtn, Exception):
|
||||||
|
raise self.rtn
|
||||||
|
else:
|
||||||
|
return self.rtn
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self.calls = 0
|
||||||
|
|
||||||
|
one = AwhileMaker(0.2, {"value": 2})
|
||||||
|
two = AwhileMaker(0.6, ValueError("I'm not good"))
|
||||||
|
builder = StateGraph(State)
|
||||||
|
builder.add_node("one", one)
|
||||||
|
builder.add_node("two", two)
|
||||||
|
builder.add_edge(START, "one")
|
||||||
|
builder.add_edge(START, "two")
|
||||||
|
graph = builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
thread1: RunnableConfig = {"configurable": {"thread_id": 1}}
|
||||||
|
with pytest.raises(ValueError, match="I'm not good"):
|
||||||
|
graph.invoke({"value": 1}, thread1)
|
||||||
|
|
||||||
|
# both nodes should have been called once
|
||||||
|
assert one.calls == 1
|
||||||
|
assert two.calls == 1
|
||||||
|
|
||||||
|
# latest checkpoint should be before nodes "one", "two"
|
||||||
|
state = graph.get_state(thread1)
|
||||||
|
assert state is not None
|
||||||
|
assert state.values == {"value": 1}
|
||||||
|
assert state.next == ("one", "two")
|
||||||
|
assert state.metadata == {"source": "loop", "step": 0, "writes": None}
|
||||||
|
# should contain pending write of "one"
|
||||||
|
checkpoint = checkpointer.get_tuple(thread1)
|
||||||
|
assert checkpoint is not None
|
||||||
|
assert checkpoint.pending_writes == [
|
||||||
|
(AnyStr(), "one", "one"),
|
||||||
|
(AnyStr(), "value", 2),
|
||||||
|
]
|
||||||
|
# both pending writes come from same task
|
||||||
|
assert checkpoint.pending_writes[0][0] == checkpoint.pending_writes[1][0]
|
||||||
|
|
||||||
|
# resume execution
|
||||||
|
with pytest.raises(ValueError, match="I'm not good"):
|
||||||
|
graph.invoke(None, thread1)
|
||||||
|
|
||||||
|
# node "one" succeeded previously, so shouldn't be called again
|
||||||
|
assert one.calls == 1
|
||||||
|
# node "two" should have been called once again
|
||||||
|
assert two.calls == 2
|
||||||
|
|
||||||
|
# confirm no new checkpoints saved
|
||||||
|
state_two = graph.get_state(thread1)
|
||||||
|
assert state_two == state
|
||||||
|
|
||||||
|
# resume execution, without exception
|
||||||
|
two.rtn = {"value": 3}
|
||||||
|
# both the pending write and the new write were applied, 1 + 2 + 3 = 6
|
||||||
|
assert graph.invoke(None, thread1) == {"value": 6}
|
||||||
|
finally:
|
||||||
|
if getattr(checkpointer, "__exit__", None):
|
||||||
|
checkpointer.__exit__(None, None, None)
|
||||||
|
|
||||||
|
|
||||||
|
def test_cond_edge_after_send() -> None:
|
||||||
|
class Node:
|
||||||
|
def __init__(self, name: str):
|
||||||
|
self.name = name
|
||||||
|
setattr(self, "__name__", name)
|
||||||
|
|
||||||
|
def __call__(self, state):
|
||||||
|
return state + [self.name]
|
||||||
|
|
||||||
|
def send_for_fun(state):
|
||||||
|
return [Send("2", state)]
|
||||||
|
|
||||||
|
def route_to_three(state) -> Literal["3"]:
|
||||||
|
return "3"
|
||||||
|
|
||||||
|
builder = StateGraph(list)
|
||||||
|
builder.add_node(Node("1"))
|
||||||
|
builder.add_node(Node("2"))
|
||||||
|
builder.add_node(Node("3"))
|
||||||
|
builder.add_edge(START, "1")
|
||||||
|
builder.add_conditional_edges("1", send_for_fun)
|
||||||
|
builder.add_conditional_edges("2", route_to_three)
|
||||||
|
graph = builder.compile()
|
||||||
|
|
||||||
|
assert graph.invoke(["0"]) == ["0", "1", "2", "3"]
|
||||||
|
|
||||||
|
|
||||||
def test_invoke_checkpoint_sqlite(mocker: MockerFixture) -> None:
|
def test_invoke_checkpoint_sqlite(mocker: MockerFixture) -> None:
|
||||||
adder = mocker.Mock(side_effect=lambda x: x["total"] + x["input"])
|
adder = mocker.Mock(side_effect=lambda x: x["total"] + x["input"])
|
||||||
|
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import operator
|
import operator
|
||||||
import time
|
|
||||||
from collections import Counter
|
from collections import Counter
|
||||||
from contextlib import asynccontextmanager, contextmanager
|
from contextlib import asynccontextmanager, contextmanager
|
||||||
from typing import (
|
from typing import (
|
||||||
@@ -11,8 +10,11 @@ from typing import (
|
|||||||
AsyncIterator,
|
AsyncIterator,
|
||||||
Dict,
|
Dict,
|
||||||
Generator,
|
Generator,
|
||||||
|
List,
|
||||||
|
Literal,
|
||||||
Optional,
|
Optional,
|
||||||
Sequence,
|
Sequence,
|
||||||
|
Tuple,
|
||||||
TypedDict,
|
TypedDict,
|
||||||
Union,
|
Union,
|
||||||
)
|
)
|
||||||
@@ -75,6 +77,12 @@ async def test_checkpoint_errors() -> None:
|
|||||||
) -> RunnableConfig:
|
) -> RunnableConfig:
|
||||||
raise ValueError("Faulty put")
|
raise ValueError("Faulty put")
|
||||||
|
|
||||||
|
class FaultyPutWritesCheckpointer(MemorySaver):
|
||||||
|
async def aput_writes(
|
||||||
|
self, config: RunnableConfig, writes: List[Tuple[str, Any]], task_id: str
|
||||||
|
) -> RunnableConfig:
|
||||||
|
raise ValueError("Faulty put_writes")
|
||||||
|
|
||||||
class FaultyVersionCheckpointer(MemorySaver):
|
class FaultyVersionCheckpointer(MemorySaver):
|
||||||
def get_next_version(self, current: Optional[int], channel: BaseChannel) -> int:
|
def get_next_version(self, current: Optional[int], channel: BaseChannel) -> int:
|
||||||
raise ValueError("Faulty get_next_version")
|
raise ValueError("Faulty get_next_version")
|
||||||
@@ -82,10 +90,9 @@ async def test_checkpoint_errors() -> None:
|
|||||||
def logic(inp: str) -> str:
|
def logic(inp: str) -> str:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
builder = Graph()
|
builder = StateGraph(Annotated[str, operator.add])
|
||||||
builder.add_node("agent", logic)
|
builder.add_node("agent", logic)
|
||||||
builder.set_entry_point("agent")
|
builder.add_edge(START, "agent")
|
||||||
builder.set_finish_point("agent")
|
|
||||||
|
|
||||||
graph = builder.compile(checkpointer=FaultyGetCheckpointer())
|
graph = builder.compile(checkpointer=FaultyGetCheckpointer())
|
||||||
with pytest.raises(ValueError, match="Faulty get_tuple"):
|
with pytest.raises(ValueError, match="Faulty get_tuple"):
|
||||||
@@ -123,6 +130,21 @@ async def test_checkpoint_errors() -> None:
|
|||||||
):
|
):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
# add a parallel node
|
||||||
|
builder.add_node("parallel", logic)
|
||||||
|
builder.add_edge(START, "parallel")
|
||||||
|
graph = builder.compile(checkpointer=FaultyPutWritesCheckpointer())
|
||||||
|
with pytest.raises(ValueError, match="Faulty put_writes"):
|
||||||
|
await graph.ainvoke("", {"configurable": {"thread_id": "thread-1"}})
|
||||||
|
with pytest.raises(ValueError, match="Faulty put_writes"):
|
||||||
|
async for _ in graph.astream("", {"configurable": {"thread_id": "thread-2"}}):
|
||||||
|
pass
|
||||||
|
with pytest.raises(ValueError, match="Faulty put_writes"):
|
||||||
|
async for _ in graph.astream_events(
|
||||||
|
"", {"configurable": {"thread_id": "thread-3"}}, version="v2"
|
||||||
|
):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
async def test_node_cancellation_on_external_cancel() -> None:
|
async def test_node_cancellation_on_external_cancel() -> None:
|
||||||
inner_task_cancelled = False
|
inner_task_cancelled = False
|
||||||
@@ -213,6 +235,11 @@ async def test_step_timeout_on_stream_hang() -> None:
|
|||||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||||
None,
|
None,
|
||||||
],
|
],
|
||||||
|
ids=[
|
||||||
|
"memory",
|
||||||
|
"aiosqlite",
|
||||||
|
"none",
|
||||||
|
],
|
||||||
)
|
)
|
||||||
async def test_cancel_graph_astream(
|
async def test_cancel_graph_astream(
|
||||||
checkpointer: Optional[BaseCheckpointSaver],
|
checkpointer: Optional[BaseCheckpointSaver],
|
||||||
@@ -279,6 +306,11 @@ async def test_cancel_graph_astream(
|
|||||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||||
None,
|
None,
|
||||||
],
|
],
|
||||||
|
ids=[
|
||||||
|
"memory",
|
||||||
|
"aiosqlite",
|
||||||
|
"none",
|
||||||
|
],
|
||||||
)
|
)
|
||||||
async def test_cancel_graph_astream_events_v2(
|
async def test_cancel_graph_astream_events_v2(
|
||||||
checkpointer: Optional[BaseCheckpointSaver],
|
checkpointer: Optional[BaseCheckpointSaver],
|
||||||
@@ -327,7 +359,6 @@ async def test_cancel_graph_astream_events_v2(
|
|||||||
) as stream:
|
) as stream:
|
||||||
async for chunk in stream:
|
async for chunk in stream:
|
||||||
if chunk["event"] == "on_chain_stream" and not chunk["parent_ids"]:
|
if chunk["event"] == "on_chain_stream" and not chunk["parent_ids"]:
|
||||||
print(time.perf_counter(), "got event out here", chunk)
|
|
||||||
got_event = True
|
got_event = True
|
||||||
assert chunk["data"]["chunk"] == {"alittlewhile": {"value": 2}}
|
assert chunk["data"]["chunk"] == {"alittlewhile": {"value": 2}}
|
||||||
break
|
break
|
||||||
@@ -1036,6 +1067,122 @@ async def test_invoke_checkpoint(mocker: MockerFixture) -> None:
|
|||||||
assert checkpoint["channel_values"].get("total") == 5
|
assert checkpoint["channel_values"].get("total") == 5
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"checkpointer",
|
||||||
|
[
|
||||||
|
MemorySaverAssertImmutable(),
|
||||||
|
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||||
|
],
|
||||||
|
ids=[
|
||||||
|
"memory",
|
||||||
|
"sqlite",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_pending_writes_resume(checkpointer: BaseCheckpointSaver) -> None:
|
||||||
|
try:
|
||||||
|
|
||||||
|
class State(TypedDict):
|
||||||
|
value: Annotated[int, operator.add]
|
||||||
|
|
||||||
|
class AwhileMaker:
|
||||||
|
def __init__(self, sleep: float, rtn: Union[Dict, Exception]) -> None:
|
||||||
|
self.sleep = sleep
|
||||||
|
self.rtn = rtn
|
||||||
|
self.reset()
|
||||||
|
|
||||||
|
async def __call__(self, input: State) -> Any:
|
||||||
|
self.calls += 1
|
||||||
|
await asyncio.sleep(self.sleep)
|
||||||
|
if isinstance(self.rtn, Exception):
|
||||||
|
raise self.rtn
|
||||||
|
else:
|
||||||
|
return self.rtn
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self.calls = 0
|
||||||
|
|
||||||
|
one = AwhileMaker(0.2, {"value": 2})
|
||||||
|
two = AwhileMaker(0.6, ValueError("I'm not good"))
|
||||||
|
builder = StateGraph(State)
|
||||||
|
builder.add_node("one", one)
|
||||||
|
builder.add_node("two", two)
|
||||||
|
builder.add_edge(START, "one")
|
||||||
|
builder.add_edge(START, "two")
|
||||||
|
graph = builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
thread1: RunnableConfig = {"configurable": {"thread_id": 1}}
|
||||||
|
with pytest.raises(ValueError, match="I'm not good"):
|
||||||
|
await graph.ainvoke({"value": 1}, thread1)
|
||||||
|
|
||||||
|
# both nodes should have been called once
|
||||||
|
assert one.calls == 1
|
||||||
|
assert two.calls == 1
|
||||||
|
|
||||||
|
# latest checkpoint should be before nodes "one", "two"
|
||||||
|
state = await graph.aget_state(thread1)
|
||||||
|
assert state is not None
|
||||||
|
assert state.values == {"value": 1}
|
||||||
|
assert state.next == ("one", "two")
|
||||||
|
assert state.metadata == {"source": "loop", "step": 0, "writes": None}
|
||||||
|
# should contain pending write of "one"
|
||||||
|
checkpoint = await checkpointer.aget_tuple(thread1)
|
||||||
|
assert checkpoint is not None
|
||||||
|
assert checkpoint.pending_writes == [
|
||||||
|
(AnyStr(), "one", "one"),
|
||||||
|
(AnyStr(), "value", 2),
|
||||||
|
]
|
||||||
|
# both pending writes come from same task
|
||||||
|
assert checkpoint.pending_writes[0][0] == checkpoint.pending_writes[1][0]
|
||||||
|
|
||||||
|
# resume execution
|
||||||
|
with pytest.raises(ValueError, match="I'm not good"):
|
||||||
|
await graph.ainvoke(None, thread1)
|
||||||
|
|
||||||
|
# node "one" succeeded previously, so shouldn't be called again
|
||||||
|
assert one.calls == 1
|
||||||
|
# node "two" should have been called once again
|
||||||
|
assert two.calls == 2
|
||||||
|
|
||||||
|
# confirm no new checkpoints saved
|
||||||
|
state_two = await graph.aget_state(thread1)
|
||||||
|
assert state_two == state
|
||||||
|
|
||||||
|
# resume execution, without exception
|
||||||
|
two.rtn = {"value": 3}
|
||||||
|
# both the pending write and the new write were applied, 1 + 2 + 3 = 6
|
||||||
|
assert await graph.ainvoke(None, thread1) == {"value": 6}
|
||||||
|
finally:
|
||||||
|
if getattr(checkpointer, "__aexit__", None):
|
||||||
|
await checkpointer.__aexit__(None, None, None)
|
||||||
|
|
||||||
|
|
||||||
|
async def test_cond_edge_after_send() -> None:
|
||||||
|
class Node:
|
||||||
|
def __init__(self, name: str):
|
||||||
|
self.name = name
|
||||||
|
setattr(self, "__name__", name)
|
||||||
|
|
||||||
|
async def __call__(self, state):
|
||||||
|
return state + [self.name]
|
||||||
|
|
||||||
|
async def send_for_fun(state):
|
||||||
|
return [Send("2", state)]
|
||||||
|
|
||||||
|
async def route_to_three(state) -> Literal["3"]:
|
||||||
|
return "3"
|
||||||
|
|
||||||
|
builder = StateGraph(list)
|
||||||
|
builder.add_node(Node("1"))
|
||||||
|
builder.add_node(Node("2"))
|
||||||
|
builder.add_node(Node("3"))
|
||||||
|
builder.add_edge(START, "1")
|
||||||
|
builder.add_conditional_edges("1", send_for_fun)
|
||||||
|
builder.add_conditional_edges("2", route_to_three)
|
||||||
|
graph = builder.compile()
|
||||||
|
|
||||||
|
assert await graph.ainvoke(["0"]) == ["0", "1", "2", "3"]
|
||||||
|
|
||||||
|
|
||||||
async def test_invoke_checkpoint_aiosqlite(mocker: MockerFixture) -> None:
|
async def test_invoke_checkpoint_aiosqlite(mocker: MockerFixture) -> None:
|
||||||
add_one = mocker.Mock(side_effect=lambda x: x["total"] + x["input"])
|
add_one = mocker.Mock(side_effect=lambda x: x["total"] + x["input"])
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user