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:
Nuno Campos
2024-07-10 14:28:06 -07:00
committed by GitHub
parent 4d2456be40
commit dfb2ac321f
10 changed files with 614 additions and 124 deletions
@@ -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."""
+44 -1
View File
@@ -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
)
+66 -34
View File
@@ -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.
+101 -29
View File
@@ -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:
+3 -3
View File
@@ -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
) )
) )
+2 -2
View File
@@ -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)
+1
View File
@@ -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):
+133 -3
View File
@@ -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"])
+152 -5
View File
@@ -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"])