diff --git a/langgraph/checkpoint/aiosqlite.py b/langgraph/checkpoint/aiosqlite.py index 81a5a0943..ada5c10d1 100644 --- a/langgraph/checkpoint/aiosqlite.py +++ b/langgraph/checkpoint/aiosqlite.py @@ -44,6 +44,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): CREATE TABLE IF NOT EXISTS checkpoints ( thread_id TEXT NOT NULL, thread_ts TEXT NOT NULL, + parent_ts TEXT, checkpoint BLOB, PRIMARY KEY (thread_id, thread_ts) ); @@ -57,17 +58,28 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): await self.setup() if config["configurable"].get("thread_ts"): async with self.conn.execute( - "SELECT checkpoint FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", + "SELECT checkpoint, parent_ts FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", ( config["configurable"]["thread_id"], config["configurable"]["thread_ts"], ), ) as cursor: if value := await cursor.fetchone(): - return CheckpointTuple(config, pickle.loads(value[0])) + return CheckpointTuple( + config, + pickle.loads(value[0]), + { + "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, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", + "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", (config["configurable"]["thread_id"],), ) as cursor: if value := await cursor.fetchone(): @@ -78,19 +90,30 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): "thread_ts": value[1], } }, - pickle.loads(value[2]), + pickle.loads(value[3]), + { + "configurable": { + "thread_id": value[0], + "thread_ts": value[2], + } + } + if value[2] + else None, ) async def alist(self, config: RunnableConfig) -> AsyncIterator[CheckpointTuple]: await self.setup() async with self.conn.execute( - "SELECT thread_id, thread_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC", + "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC", (config["configurable"]["thread_id"],), ) as cursor: - async for thread_id, thread_ts, value in cursor: + async for thread_id, thread_ts, parent_ts, value in cursor: yield CheckpointTuple( {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, pickle.loads(value), + {"configurable": {"thread_id": thread_id, "thread_ts": parent_ts}} + if parent_ts + else None, ) async def aput( @@ -98,10 +121,11 @@ class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager): ) -> RunnableConfig: await self.setup() async with self.conn.execute( - "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, checkpoint) VALUES (?, ?, ?)", + "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint) VALUES (?, ?, ?, ?)", ( config["configurable"]["thread_id"], checkpoint["ts"], + config["configurable"].get("thread_ts"), pickle.dumps(checkpoint), ), ): diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index b5000a303..7fe967a8b 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -52,6 +52,7 @@ class CheckpointAt(StrEnum): class CheckpointTuple(NamedTuple): config: RunnableConfig checkpoint: Checkpoint + parent_config: Optional[RunnableConfig] = None CheckpointThreadId = ConfigurableFieldSpec( diff --git a/langgraph/checkpoint/sqlite.py b/langgraph/checkpoint/sqlite.py index b7dbc04c8..3a292fb5a 100644 --- a/langgraph/checkpoint/sqlite.py +++ b/langgraph/checkpoint/sqlite.py @@ -43,6 +43,7 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): CREATE TABLE IF NOT EXISTS checkpoints ( thread_id TEXT NOT NULL, thread_ts TEXT NOT NULL, + parent_ts TEXT, checkpoint BLOB, PRIMARY KEY (thread_id, thread_ts) ); @@ -66,17 +67,28 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): with self.cursor(transaction=False) as cur: if config["configurable"].get("thread_ts"): cur.execute( - "SELECT checkpoint FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", + "SELECT checkpoint, parent_ts FROM checkpoints WHERE thread_id = ? AND thread_ts = ?", ( config["configurable"]["thread_id"], config["configurable"]["thread_ts"], ), ) if value := cur.fetchone(): - return CheckpointTuple(config, pickle.loads(value[0])) + return CheckpointTuple( + config, + pickle.loads(value[0]), + { + "configurable": { + "thread_id": config["configurable"]["thread_id"], + "thread_ts": value[1], + } + } + if value[1] + else None, + ) else: cur.execute( - "SELECT thread_id, thread_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", + "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC LIMIT 1", (config["configurable"]["thread_id"],), ) if value := cur.fetchone(): @@ -87,28 +99,45 @@ class SqliteSaver(BaseCheckpointSaver, AbstractContextManager): "thread_ts": value[1], } }, - pickle.loads(value[2]), + pickle.loads(value[3]), + { + "configurable": { + "thread_id": value[0], + "thread_ts": value[2], + } + } + if value[2] + else None, ) def list(self, config: RunnableConfig) -> Iterator[CheckpointTuple]: with self.cursor(transaction=False) as cur: cur.execute( - "SELECT thread_id, thread_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC", + "SELECT thread_id, thread_ts, parent_ts, checkpoint FROM checkpoints WHERE thread_id = ? ORDER BY thread_ts DESC", (config["configurable"]["thread_id"],), ) - for thread_id, thread_ts, value in cur: + for thread_id, thread_ts, parent_ts, value in cur: yield CheckpointTuple( {"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}}, pickle.loads(value), + { + "configurable": { + "thread_id": thread_id, + "thread_ts": parent_ts, + } + } + if parent_ts + else None, ) def put(self, config: RunnableConfig, checkpoint: Checkpoint) -> RunnableConfig: with self.cursor() as cur: cur.execute( - "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, checkpoint) VALUES (?, ?, ?)", + "INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, parent_ts, checkpoint) VALUES (?, ?, ?, ?)", ( config["configurable"]["thread_id"], checkpoint["ts"], + config["configurable"].get("thread_ts"), pickle.dumps(checkpoint), ), ) diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index 1a6bc6f7c..e7a45a38b 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -165,6 +165,8 @@ class StateSnapshot(NamedTuple): """Nodes to execute in the next step, if any""" config: RunnableConfig """Config used to fetch this snapshot""" + parent_config: Optional[RunnableConfig] = None + """Config used to fetch the parent snapshot, if any""" class Pregel( @@ -328,7 +330,7 @@ class Pregel( if not self.checkpointer: raise ValueError("No checkpointer set") - for config, checkpoint in self.checkpointer.list(config): + for config, checkpoint, parent_config in self.checkpointer.list(config): with ChannelsManager(self.channels, checkpoint) as channels: _, next_tasks = _prepare_next_tasks( checkpoint, self.nodes, channels, update_seen=False @@ -344,6 +346,7 @@ class Pregel( else values, tuple(name for _, _, name in next_tasks), config, + parent_config, ) async def aget_state_history( @@ -352,7 +355,7 @@ class Pregel( if not self.checkpointer: raise ValueError("No checkpointer set") - async for config, checkpoint in self.checkpointer.alist(config): + async for config, checkpoint, parent_config in self.checkpointer.alist(config): async with AsyncChannelsManager(self.channels, checkpoint) as channels: _, next_tasks = _prepare_next_tasks( checkpoint, self.nodes, channels, update_seen=False @@ -368,6 +371,7 @@ class Pregel( else values, tuple(name for _, _, name in next_tasks), config, + parent_config, ) def update_state( @@ -473,7 +477,10 @@ class Pregel( # copy nodes to ignore mutations during execution processes = {**self.nodes} # get checkpoint from saver, or create an empty one - checkpoint = self.checkpointer.get(config) if self.checkpointer else None + checkpoint_config = config + checkpoint = ( + self.checkpointer.get(checkpoint_config) if self.checkpointer else None + ) checkpoint = checkpoint or empty_checkpoint() # create channels from checkpoint with ChannelsManager( @@ -595,7 +602,9 @@ class Pregel( or interrupt_before_nodes ): checkpoint = create_checkpoint(checkpoint, channels) - self.checkpointer.put(config, checkpoint) + checkpoint_config = self.checkpointer.put( + checkpoint_config, checkpoint + ) # with this step's checkpoint, if _should_interrupt( @@ -613,7 +622,7 @@ class Pregel( and not interrupt_before_nodes ): checkpoint = create_checkpoint(checkpoint, channels) - self.checkpointer.put(config, checkpoint) + self.checkpointer.put(checkpoint_config, checkpoint) finally: # cancel any pending tasks when generator is interrupted try: @@ -652,8 +661,11 @@ class Pregel( # copy nodes to ignore mutations during execution processes = {**self.nodes} # get checkpoint from saver, or create an empty one + checkpoint_config = config checkpoint = ( - await self.checkpointer.aget(config) if self.checkpointer else None + await self.checkpointer.aget(checkpoint_config) + if self.checkpointer + else None ) checkpoint = checkpoint or empty_checkpoint() # create channels from checkpoint @@ -781,7 +793,9 @@ class Pregel( or interrupt_before_nodes ): checkpoint = create_checkpoint(checkpoint, channels) - await self.checkpointer.aput(config, checkpoint) + checkpoint_config = await self.checkpointer.aput( + checkpoint_config, checkpoint + ) # with this step's checkpoint if _should_interrupt( @@ -799,7 +813,7 @@ class Pregel( and not interrupt_before_nodes ): checkpoint = create_checkpoint(checkpoint, channels) - await self.checkpointer.aput(config, checkpoint) + await self.checkpointer.aput(checkpoint_config, checkpoint) finally: # cancel any pending tasks when generator is interrupted try: diff --git a/tests/memory_assert.py b/tests/memory_assert.py index a2d1c9ee8..50c185a12 100644 --- a/tests/memory_assert.py +++ b/tests/memory_assert.py @@ -22,4 +22,4 @@ class MemorySaverAssertImmutable(MemorySaver): checkpoint ) # call super to write checkpoint - super().put(config, checkpoint) + return super().put(config, checkpoint)