diff --git a/libs/langgraph/langgraph/checkpoint/aiosqlite.py b/libs/langgraph/langgraph/checkpoint/aiosqlite.py index d8da20de2..431684d1a 100644 --- a/libs/langgraph/langgraph/checkpoint/aiosqlite.py +++ b/libs/langgraph/langgraph/checkpoint/aiosqlite.py @@ -246,7 +246,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): # 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 <= ? ORDER BY thread_ts DESC LIMIT 1", + "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"]), diff --git a/libs/langgraph/langgraph/checkpoint/memory.py b/libs/langgraph/langgraph/checkpoint/memory.py index 2af8fc21d..b71dd0603 100644 --- a/libs/langgraph/langgraph/checkpoint/memory.py +++ b/libs/langgraph/langgraph/checkpoint/memory.py @@ -44,7 +44,7 @@ class MemorySaver(BaseCheckpointSaver): asyncio.run(coro) # Output: 2 """ - storage: defaultdict[str, dict[str, tuple[bytes, bytes]]] + storage: defaultdict[str, dict[str, tuple[bytes, bytes, Optional[str]]]] def __init__( self, @@ -70,25 +70,30 @@ class MemorySaver(BaseCheckpointSaver): Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found. """ thread_id = config["configurable"]["thread_id"] - if thread_ts := config["configurable"].get("thread_ts"): - if checkpoints := self.storage[thread_id]: - matching_keys = [key for key in checkpoints.keys() if key <= thread_ts] - ts = max(matching_keys) if matching_keys else None - 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 - ], - ) + if ts := config["configurable"].get("thread_ts"): + if saved := self.storage[thread_id].get(ts): + checkpoint, metadata, parent_ts = 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 + ], + parent_config={ + "configurable": { + "thread_id": thread_id, + "thread_ts": parent_ts, + } + } + if parent_ts + else None, + ) else: if checkpoints := self.storage[thread_id]: ts = max(checkpoints.keys()) - checkpoint, metadata = checkpoints[ts] + checkpoint, metadata, parent_ts = checkpoints[ts] writes = self.writes[(thread_id, ts)] return CheckpointTuple( config={"configurable": {"thread_id": thread_id, "thread_ts": ts}}, @@ -97,6 +102,14 @@ class MemorySaver(BaseCheckpointSaver): pending_writes=[ (id, c, self.serde.loads(v)) for id, c, v in writes ], + parent_config={ + "configurable": { + "thread_id": thread_id, + "thread_ts": parent_ts, + } + } + if parent_ts + else None, ) def list( @@ -122,7 +135,7 @@ class MemorySaver(BaseCheckpointSaver): """ thread_ids = (config["configurable"]["thread_id"],) if config else self.storage for thread_id in thread_ids: - for ts, (checkpoint, metadata_b) in sorted( + for ts, (checkpoint, metadata_b, parent_ts) in sorted( self.storage[thread_id].items(), key=lambda x: x[0], reverse=True ): # filter by thread_ts @@ -147,6 +160,14 @@ class MemorySaver(BaseCheckpointSaver): config={"configurable": {"thread_id": thread_id, "thread_ts": ts}}, checkpoint=self.serde.loads(checkpoint), metadata=metadata, + parent_config={ + "configurable": { + "thread_id": thread_id, + "thread_ts": parent_ts, + } + } + if parent_ts + else None, ) def put( @@ -172,6 +193,7 @@ class MemorySaver(BaseCheckpointSaver): checkpoint["id"]: ( self.serde.dumps(checkpoint), self.serde.dumps(metadata), + config["configurable"].get("thread_ts"), # parent ) } ) diff --git a/libs/langgraph/langgraph/checkpoint/sqlite.py b/libs/langgraph/langgraph/checkpoint/sqlite.py index 737d426b7..6ac9ee593 100644 --- a/libs/langgraph/langgraph/checkpoint/sqlite.py +++ b/libs/langgraph/langgraph/checkpoint/sqlite.py @@ -245,7 +245,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): # find the latest checkpoint for the thread_id if config["configurable"].get("thread_ts"): cur.execute( - "SELECT thread_id, thread_ts, parent_ts, checkpoint, metadata FROM checkpoints WHERE thread_id = ? AND thread_ts <= ? ORDER BY thread_ts DESC LIMIT 1", + "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"]), diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 5c288aa50..13ce7e183 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -382,6 +382,7 @@ def prepare_next_tasks( CONFIG_KEY_CHECKPOINTER: checkpointer, CONFIG_KEY_RESUMING: is_resuming, "thread_id": thread_id, + "thread_ts": checkpoint["id"], }, ), triggers, diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index 389c59e81..70c1e69b7 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -70,14 +70,15 @@ def map_debug_tasks( if config is not None and TAG_HIDDEN in config.get("tags", []): continue + metadata = config["metadata"].copy() + metadata.pop("thread_ts", None) + yield { "type": "task", "timestamp": ts, "step": step, "payload": { - "id": str( - uuid5(TASK_NAMESPACE, json.dumps((name, step, config["metadata"]))) - ), + "id": str(uuid5(TASK_NAMESPACE, json.dumps((name, step, metadata)))), "name": name, "input": input, "triggers": triggers, @@ -95,14 +96,15 @@ def map_debug_task_results( if config is not None and TAG_HIDDEN in config.get("tags", []): continue + metadata = config["metadata"].copy() + metadata.pop("thread_ts", None) + yield { "type": "task_result", "timestamp": ts, "step": step, "payload": { - "id": str( - uuid5(TASK_NAMESPACE, json.dumps((name, step, config["metadata"]))) - ), + "id": str(uuid5(TASK_NAMESPACE, json.dumps((name, step, metadata)))), "name": name, "result": [w for w in writes if w[0] in stream_channels_list], }, diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index a3e29564d..d2abf8372 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -1,5 +1,4 @@ import asyncio -import concurrent.futures from collections import deque from contextlib import AsyncExitStack, ExitStack from types import TracebackType @@ -69,6 +68,7 @@ if TYPE_CHECKING: V = TypeVar("V") INPUT_DONE = object() INPUT_RESUMING = object() +EMPTY_LIST = [] class PregelLoop: @@ -126,9 +126,9 @@ class PregelLoop: def tick( self, *, - output_keys: Union[str, Sequence[str]] = None, - interrupt_after: Optional[Sequence[str]] = None, - interrupt_before: Optional[Sequence[str]] = None, + output_keys: Union[str, Sequence[str]] = EMPTY_LIST, + interrupt_after: Sequence[str] = EMPTY_LIST, + interrupt_before: Sequence[str] = EMPTY_LIST, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, ) -> bool: """Execute a single iteration of the Pregel loop. @@ -208,7 +208,12 @@ class PregelLoop: # if all tasks have finished, re-tick if all(task.writes for task in self.tasks): - return self.tick() + return self.tick( + output_keys=output_keys, + interrupt_after=interrupt_after, + interrupt_before=interrupt_before, + manager=manager, + ) # before execution, check if we should interrupt if should_interrupt(self.checkpoint, interrupt_before, self.tasks): @@ -267,10 +272,7 @@ class PregelLoop: # done with input self.input = INPUT_RESUMING if is_resuming else INPUT_DONE - def _put_checkpoint( - self, - metadata: CheckpointMetadata, - ) -> concurrent.futures.Future: + def _put_checkpoint(self, metadata: CheckpointMetadata) -> None: # assign step metadata["step"] = self.step # bail if no checkpointer @@ -278,10 +280,17 @@ class PregelLoop: # create new checkpoint self.checkpoint_metadata = metadata self.checkpoint = create_checkpoint( - self.checkpoint, self.channels, self.step + self.checkpoint, + self.channels, + self.step, + # child graphs keep at most one checkpoint per parent checkpoint + # this is achieved by writing child checkpoints as progress is made + # (so that error recovery / resuming from interrupt don't lose work) + # but doing so always with an id equal to that of the parent checkpoint + id=self.config["configurable"]["thread_ts"] if self.is_nested else None, ) # save it, without blocking - fut = self.submit( + self.submit( self.checkpointer_put, self.checkpoint_config, copy_checkpoint(self.checkpoint), @@ -305,12 +314,8 @@ class PregelLoop: self.checkpoint_metadata, ) ) - else: - fut = concurrent.futures.Future() - fut.set_result(None) # increment step self.step += 1 - return fut class SyncPregelLoop(PregelLoop, ContextManager): @@ -380,9 +385,6 @@ class SyncPregelLoop(PregelLoop, ContextManager): if exc_value.args[0] is self: # interrupt raised by this loop exc_value.args = (object(),) - else: - # interrupt raised by a nested loop, save interrupt checkpoint - self._put_checkpoint({"source": "interrupt"}).result() if not self.is_nested: # in outer graph, catch interrupt del self.graph @@ -464,9 +466,6 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): if exc_value.args[0] is self: # interrupt raised by this loop exc_value.args = (object(),) - else: - # interrupt raised by a nested loop, save interrupt checkpoint - self._put_checkpoint({"source": "interrupt"}) if not self.is_nested: # in outer graph, catch interrupt del self.graph diff --git a/libs/langgraph/tests/memory_assert.py b/libs/langgraph/tests/memory_assert.py index db1212ae9..6709a404a 100644 --- a/libs/langgraph/tests/memory_assert.py +++ b/libs/langgraph/tests/memory_assert.py @@ -85,7 +85,7 @@ class MemorySaverAssertCheckpointMetadata(MemorySaver): configurable = config["configurable"].copy() # remove thread_ts to make testing simpler - configurable.pop("thread_ts", None) + thread_ts = configurable.pop("thread_ts", None) self.storage[config["configurable"]["thread_id"]].update( { @@ -93,6 +93,7 @@ class MemorySaverAssertCheckpointMetadata(MemorySaver): self.serde.dumps(checkpoint), # merge configurable fields and metadata self.serde.dumps({**configurable, **metadata}), + thread_ts, ) } )