From dfb2ac321f00a6e12f3601fce311a37ca6d620ce Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 10 Jul 2024 14:28:06 -0700 Subject: [PATCH] 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 --- .../langgraph/checkpoint/aiosqlite.py | 137 +++++++++------ libs/langgraph/langgraph/checkpoint/base.py | 22 +++ libs/langgraph/langgraph/checkpoint/memory.py | 45 ++++- libs/langgraph/langgraph/checkpoint/sqlite.py | 100 +++++++---- libs/langgraph/langgraph/pregel/__init__.py | 130 +++++++++++---- libs/langgraph/langgraph/pregel/debug.py | 6 +- libs/langgraph/langgraph/pregel/io.py | 4 +- libs/langgraph/langgraph/pregel/types.py | 1 + libs/langgraph/tests/test_pregel.py | 136 ++++++++++++++- libs/langgraph/tests/test_pregel_async.py | 157 +++++++++++++++++- 10 files changed, 614 insertions(+), 124 deletions(-) diff --git a/libs/langgraph/langgraph/checkpoint/aiosqlite.py b/libs/langgraph/langgraph/checkpoint/aiosqlite.py index b636e02d3..431684d1a 100644 --- a/libs/langgraph/langgraph/checkpoint/aiosqlite.py +++ b/libs/langgraph/langgraph/checkpoint/aiosqlite.py @@ -2,7 +2,16 @@ import asyncio import functools from contextlib import AbstractAsyncContextManager 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 from langchain_core.runnables import RunnableConfig @@ -203,6 +212,15 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): metadata BLOB, 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() @@ -224,56 +242,58 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found. """ await self.setup() - if config["configurable"].get("thread_ts"): - async with self.conn.execute( - "SELECT checkpoint, parent_ts, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", - ( - str(config["configurable"]["thread_id"]), - str(config["configurable"]["thread_ts"]), - ), - ) as cursor: - if value := await cursor.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: - async with self.conn.execute( - "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", - (str(config["configurable"]["thread_id"]),), - ) as cursor: - if value := await cursor.fetchone(): - return CheckpointTuple( + async with self.conn.cursor() as cur: + # find the latest checkpoint for the thread_id + if config["configurable"].get("thread_ts"): + await cur.execute( + "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_ts"]), + ), + ) + else: + await cur.execute( + "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", + (str(config["configurable"]["thread_id"]),), + ) + # if a checkpoint is found, return it + if value := await cur.fetchone(): + if not config["configurable"].get("thread_ts"): + config = { + "configurable": { + "thread_id": value[0], + "thread_ts": value[1], + } + } + # find any pending writes + await 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": { "thread_id": value[0], - "thread_ts": value[1], + "thread_ts": value[2], } - }, - self.serde.loads(value[3]), - self.serde.loads(value[4]) if value[4] is not None else {}, - ( - { - "configurable": { - "thread_id": value[0], - "thread_ts": value[2], - } - } - if value[2] - else None - ), - ) + } + if value[2] + else None + ), + [ + (task_id, channel, self.serde.loads(value)) + async for task_id, channel, value in cur + ], + ) async def alist( self, @@ -358,3 +378,26 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): "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() diff --git a/libs/langgraph/langgraph/checkpoint/base.py b/libs/langgraph/langgraph/checkpoint/base.py index fab38fa99..7fa64b241 100644 --- a/libs/langgraph/langgraph/checkpoint/base.py +++ b/libs/langgraph/langgraph/checkpoint/base.py @@ -10,6 +10,7 @@ from typing import ( Literal, NamedTuple, Optional, + Tuple, TypedDict, TypeVar, Union, @@ -117,6 +118,7 @@ class CheckpointTuple(NamedTuple): checkpoint: Checkpoint metadata: CheckpointMetadata parent_config: Optional[RunnableConfig] = None + pending_writes: Optional[List[Tuple[str, str, Any]]] = None CheckpointThreadId = ConfigurableFieldSpec( @@ -177,6 +179,16 @@ class BaseCheckpointSaver(ABC): ) -> RunnableConfig: 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]: if value := await self.aget_tuple(config): return value.checkpoint @@ -203,6 +215,16 @@ class BaseCheckpointSaver(ABC): ) -> RunnableConfig: 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: """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.""" diff --git a/libs/langgraph/langgraph/checkpoint/memory.py b/libs/langgraph/langgraph/checkpoint/memory.py index 5f87aa736..bd5dc6fd1 100644 --- a/libs/langgraph/langgraph/checkpoint/memory.py +++ b/libs/langgraph/langgraph/checkpoint/memory.py @@ -1,7 +1,7 @@ import asyncio from collections import defaultdict 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 @@ -53,6 +53,7 @@ class MemorySaver(BaseCheckpointSaver): ) -> None: super().__init__(serde=serde) self.storage = defaultdict(dict) + self.writes = defaultdict(list) def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: """Get a checkpoint tuple from the in-memory storage. @@ -72,19 +73,27 @@ class MemorySaver(BaseCheckpointSaver): if ts := config["configurable"].get("thread_ts"): if saved := self.storage[thread_id].get(ts): checkpoint, metadata = saved + writes = self.writes[(thread_id, ts)] return CheckpointTuple( config=config, checkpoint=self.serde.loads(checkpoint), metadata=self.serde.loads(metadata), + pending_writes=[ + (id, c, self.serde.loads(v)) for id, c, v in writes + ], ) else: if checkpoints := self.storage[thread_id]: ts = max(checkpoints.keys()) checkpoint, metadata = checkpoints[ts] + writes = self.writes[(thread_id, ts)] return CheckpointTuple( config={"configurable": {"thread_id": thread_id, "thread_ts": ts}}, checkpoint=self.serde.loads(checkpoint), metadata=self.serde.loads(metadata), + pending_writes=[ + (id, c, self.serde.loads(v)) for id, c, v in writes + ], ) 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]: """Asynchronous version of get_tuple. @@ -224,3 +257,13 @@ class MemorySaver(BaseCheckpointSaver): return await asyncio.get_running_loop().run_in_executor( 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 + ) diff --git a/libs/langgraph/langgraph/checkpoint/sqlite.py b/libs/langgraph/langgraph/checkpoint/sqlite.py index 6a4c35c85..6ac9ee593 100644 --- a/libs/langgraph/langgraph/checkpoint/sqlite.py +++ b/libs/langgraph/langgraph/checkpoint/sqlite.py @@ -171,6 +171,15 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): metadata BLOB, 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(...) """ # noqa with self.cursor(transaction=False) as cur: + # find the latest checkpoint for the thread_id if config["configurable"].get("thread_ts"): 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_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: cur.execute( "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", (str(config["configurable"]["thread_id"]),), ) - if value := cur.fetchone(): - return CheckpointTuple( + # if a checkpoint is found, return it + 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": { "thread_id": value[0], - "thread_ts": value[1], + "thread_ts": value[2], } - }, - self.serde.loads(value[3]), - self.serde.loads(value[4]) if value[4] is not None else {}, - ( - { - "configurable": { - "thread_id": value[0], - "thread_ts": value[2], - } - } - if value[2] - else None - ), - ) + } + if value[2] + else None + ), + [ + (task_id, channel, self.serde.loads(value)) + for task_id, channel, value in cur + ], + ) def list( 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]: """Get a checkpoint tuple from the database asynchronously. diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index a86674ca3..79f678fd8 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -2,6 +2,7 @@ from __future__ import annotations import asyncio import concurrent.futures +import json import time from collections import defaultdict, deque from functools import partial @@ -23,6 +24,7 @@ from typing import ( get_type_hints, overload, ) +from uuid import UUID, uuid5 from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager from langchain_core.globals import get_debug @@ -442,7 +444,7 @@ class Pregel( and signature(self.checkpointer.list).parameters.get("filter") is None ): 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 ): with ChannelsManager( @@ -489,6 +491,7 @@ class Pregel( checkpoint, metadata, parent_config, + _, ) in self.checkpointer.alist(config, before=before, limit=limit, filter=filter): async with AsyncChannelsManager( self.channels, checkpoint, config @@ -565,6 +568,7 @@ class Pregel( deque(), None, [INTERRUPT], + str(uuid5(UUID(checkpoint["id"]), INTERRUPT)), ) # execute task task.proc.invoke( @@ -653,6 +657,7 @@ class Pregel( deque(), None, [INTERRUPT], + str(uuid5(UUID(checkpoint["id"]), INTERRUPT)), ) # execute task await task.proc.ainvoke( @@ -879,6 +884,23 @@ class Pregel( self.managed_values_dict, config, self ) 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]: nonlocal checkpoint, checkpoint_config, channels @@ -963,8 +985,7 @@ class Pregel( # increment start to 0 start += 1 else: - # if received no input, take that as signal to proceed - # past previous interrupt, if any + # no input is taken as signal to proceed past previous interrupt checkpoint = copy_checkpoint(checkpoint) for k in self.stream_channels_list: 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 not next_tasks: if step == start: @@ -1027,12 +1057,15 @@ class Pregel( futures = { executor.submit(run_with_retry, task, self.retry_policy): task for task in next_tasks + if not task.writes } end_time = ( self.step_timeout + time.monotonic() if self.step_timeout else None ) + if not futures: + done, inflight = set(), set() while futures: done, inflight = concurrent.futures.wait( futures, @@ -1050,6 +1083,10 @@ class Pregel( # exception will be handled in panic_or_proceed futures.clear() 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 if "updates" in stream_modes: yield from _with_mode( @@ -1076,7 +1113,7 @@ class Pregel( # combine pending writes from all tasks pending_writes = deque[tuple[str, Any]]() - for _, _, _, writes, _, _ in next_tasks: + for _, _, _, writes, _, _, _ in next_tasks: pending_writes.extend(writes) if debug: @@ -1240,6 +1277,24 @@ class Pregel( self.managed_values_dict, config, self ) 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]: nonlocal checkpoint, checkpoint_config, channels @@ -1320,8 +1375,7 @@ class Pregel( # increment start to 0 start += 1 else: - # if received no input, take that as signal to proceed - # past previous interrupt, if any + # no input is taken as signal to proceed past previous interrupt checkpoint = copy_checkpoint(checkpoint) for k in self.stream_channels_list: 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 not next_tasks: if step == start: @@ -1387,10 +1450,13 @@ class Pregel( arun_with_retry(task, self.retry_policy, do_stream) ): task for task in next_tasks + if not task.writes } end_time = ( self.step_timeout + loop.time() if self.step_timeout else None ) + if not futures: + done, inflight = set(), set() while futures: done, inflight = await asyncio.wait( futures, @@ -1406,6 +1472,10 @@ class Pregel( # exception will be handle in panic_or_proceed futures.clear() 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 if "updates" in stream_modes: for chunk in _with_mode( @@ -1434,7 +1504,7 @@ class Pregel( # combine pending writes from all tasks pending_writes = deque[tuple[str, Any]]() - for _, _, _, writes, _, _ in next_tasks: + for _, _, _, writes, _, _, _ in next_tasks: pending_writes.extend(writes) if debug: @@ -1671,7 +1741,7 @@ def _should_interrupt( # and any triggered node is in interrupt_nodes list and any( node - for node, _, _, _, config, _ in tasks + for node, _, _, _, config, _, _ in tasks if ( (not config or TAG_HIDDEN not in config.get("tags")) if interrupt_nodes == "*" @@ -1825,6 +1895,14 @@ def _prepare_next_tasks( continue if for_execution: 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() tasks.append( PregelExecutableTask( @@ -1836,14 +1914,7 @@ def _prepare_next_tasks( merge_configs( config, processes[packet.node].config, - { - "metadata": { - "langgraph_step": step, - "langgraph_node": packet.node, - "langgraph_triggers": [TASKS], - "langgraph_task_idx": len(tasks), - } - }, + {"metadata": metadata}, ), run_name=packet.node, callbacks=( @@ -1857,11 +1928,12 @@ def _prepare_next_tasks( _local_write, writes.extend, processes, channels ), CONFIG_KEY_READ: partial( - _local_read, checkpoint, channels, tasks, config + _local_read, checkpoint, channels, writes, config ), }, ), - [TASKS], + triggers, + task_id, ) ) else: @@ -1879,7 +1951,7 @@ def _prepare_next_tasks( for name, proc in processes.items(): seen = checkpoint["versions_seen"][name] # If any of the channels read by this process were updated - if triggers := [ + if triggers := sorted( chan for chan in proc.triggers if not isinstance( @@ -1887,7 +1959,7 @@ def _prepare_next_tasks( ) and checkpoint["channel_versions"].get(chan, null_version) > seen.get(chan, null_version) - ]: + ): channels_to_consume.update(triggers) try: val = next(_proc_input(step, name, proc, managed, channels)) @@ -1906,8 +1978,14 @@ def _prepare_next_tasks( if for_execution: 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() - triggers = sorted(triggers) tasks.append( PregelExecutableTask( name, @@ -1918,14 +1996,7 @@ def _prepare_next_tasks( merge_configs( config, proc.config, - { - "metadata": { - "langgraph_step": step, - "langgraph_node": name, - "langgraph_triggers": triggers, - "langgraph_task_idx": len(tasks), - } - }, + {"metadata": metadata}, ), run_name=name, callbacks=( @@ -1948,6 +2019,7 @@ def _prepare_next_tasks( }, ), triggers, + task_id, ) ) else: diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index 7fb294903..0293690a7 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -66,7 +66,7 @@ def map_debug_tasks( step: int, tasks: list[PregelExecutableTask] ) -> Iterator[DebugOutputTask]: 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", []): continue @@ -91,7 +91,7 @@ def map_debug_task_results( stream_channels_list: Sequence[str], ) -> Iterator[DebugOutputTaskResult]: 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", []): continue @@ -138,7 +138,7 @@ def print_step_tasks(step: int, next_tasks: list[PregelExecutableTask]) -> None: ) + "\n".join( f"- {get_colored_text(name, 'green')} -> {pformat(val)}" - for name, val, _, _, _, _ in next_tasks + for name, val, _, _, _, _, _ in next_tasks ) ) diff --git a/libs/langgraph/langgraph/pregel/io.py b/libs/langgraph/langgraph/pregel/io.py index fc1fc38ea..d33741697 100644 --- a/libs/langgraph/langgraph/pregel/io.py +++ b/libs/langgraph/langgraph/pregel/io.py @@ -105,7 +105,7 @@ def map_output_updates( if isinstance(output_channels, str): if updated := [ (node, value) - for node, _, _, writes, _, _ in output_tasks + for node, _, _, writes, _, _, _ in output_tasks for chan, value in writes if chan == output_channels ]: @@ -122,7 +122,7 @@ def map_output_updates( node, {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) ]: grouped = defaultdict(list) diff --git a/libs/langgraph/langgraph/pregel/types.py b/libs/langgraph/langgraph/pregel/types.py index e47d94b57..7d23eb009 100644 --- a/libs/langgraph/langgraph/pregel/types.py +++ b/libs/langgraph/langgraph/pregel/types.py @@ -18,6 +18,7 @@ class PregelExecutableTask(NamedTuple): writes: deque[tuple[str, Any]] config: RunnableConfig triggers: list[str] + id: str class StateSnapshot(NamedTuple): diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 39b339edf..a304534f2 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -14,6 +14,7 @@ from typing import ( Literal, Optional, Sequence, + Tuple, TypedDict, Union, ) @@ -37,6 +38,7 @@ from langgraph.channels.context import Context from langgraph.channels.last_value import LastValue from langgraph.channels.topic import Topic from langgraph.checkpoint.base import ( + BaseCheckpointSaver, Checkpoint, CheckpointMetadata, CheckpointTuple, @@ -193,6 +195,12 @@ def test_checkpoint_errors() -> None: ) -> RunnableConfig: 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): def get_next_version(self, current: Optional[int], channel: BaseChannel) -> int: raise ValueError("Faulty get_next_version") @@ -200,10 +208,9 @@ def test_checkpoint_errors() -> None: def logic(inp: str) -> str: return "" - builder = Graph() + builder = StateGraph(Annotated[str, operator.add]) builder.add_node("agent", logic) - builder.set_entry_point("agent") - builder.set_finish_point("agent") + builder.add_edge(START, "agent") graph = builder.compile(checkpointer=FaultyGetCheckpointer()) 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"): 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: from langchain_core.messages import HumanMessage @@ -944,6 +958,122 @@ def test_invoke_checkpoint(mocker: MockerFixture) -> None: 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: adder = mocker.Mock(side_effect=lambda x: x["total"] + x["input"]) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index cf47292a1..41f99d53f 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -1,7 +1,6 @@ import asyncio import json import operator -import time from collections import Counter from contextlib import asynccontextmanager, contextmanager from typing import ( @@ -11,8 +10,11 @@ from typing import ( AsyncIterator, Dict, Generator, + List, + Literal, Optional, Sequence, + Tuple, TypedDict, Union, ) @@ -75,6 +77,12 @@ async def test_checkpoint_errors() -> None: ) -> RunnableConfig: 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): def get_next_version(self, current: Optional[int], channel: BaseChannel) -> int: raise ValueError("Faulty get_next_version") @@ -82,10 +90,9 @@ async def test_checkpoint_errors() -> None: def logic(inp: str) -> str: return "" - builder = Graph() + builder = StateGraph(Annotated[str, operator.add]) builder.add_node("agent", logic) - builder.set_entry_point("agent") - builder.set_finish_point("agent") + builder.add_edge(START, "agent") graph = builder.compile(checkpointer=FaultyGetCheckpointer()) with pytest.raises(ValueError, match="Faulty get_tuple"): @@ -123,6 +130,21 @@ async def test_checkpoint_errors() -> None: ): 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: inner_task_cancelled = False @@ -213,6 +235,11 @@ async def test_step_timeout_on_stream_hang() -> None: AsyncSqliteSaver.from_conn_string(":memory:"), None, ], + ids=[ + "memory", + "aiosqlite", + "none", + ], ) async def test_cancel_graph_astream( checkpointer: Optional[BaseCheckpointSaver], @@ -279,6 +306,11 @@ async def test_cancel_graph_astream( AsyncSqliteSaver.from_conn_string(":memory:"), None, ], + ids=[ + "memory", + "aiosqlite", + "none", + ], ) async def test_cancel_graph_astream_events_v2( checkpointer: Optional[BaseCheckpointSaver], @@ -327,7 +359,6 @@ async def test_cancel_graph_astream_events_v2( ) as stream: async for chunk in stream: if chunk["event"] == "on_chain_stream" and not chunk["parent_ids"]: - print(time.perf_counter(), "got event out here", chunk) got_event = True assert chunk["data"]["chunk"] == {"alittlewhile": {"value": 2}} break @@ -1036,6 +1067,122 @@ async def test_invoke_checkpoint(mocker: MockerFixture) -> None: 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: add_one = mocker.Mock(side_effect=lambda x: x["total"] + x["input"])