Add delete_thread method to Checkpointer class (#4328)

- Deletes all data associated with a thread_id
- Implemented in InMemory, Sqlite and Postgres checkpointers

Co-authored-by: Eugene Yurtsev <eyurtsev@gmail.com>
This commit is contained in:
Nuno Campos
2025-04-17 16:38:58 +00:00
committed by GitHub
co-authored by Eugene Yurtsev
parent 83bf004ad7
commit 18a9ae45f3
7 changed files with 201 additions and 4 deletions
@@ -357,6 +357,29 @@ class PostgresSaver(BasePostgresSaver):
),
)
def delete_thread(self, thread_id: str) -> None:
"""Delete all checkpoints and writes associated with a thread ID.
Args:
thread_id (str): The thread ID to delete.
Returns:
None
"""
with self._cursor(pipeline=True) as cur:
cur.execute(
"DELETE FROM checkpoints WHERE thread_id = %s",
(str(thread_id),),
)
cur.execute(
"DELETE FROM checkpoint_blobs WHERE thread_id = %s",
(str(thread_id),),
)
cur.execute(
"DELETE FROM checkpoint_writes WHERE thread_id = %s",
(str(thread_id),),
)
@contextmanager
def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]:
"""Create a database cursor as a context manager.
@@ -314,6 +314,29 @@ class AsyncPostgresSaver(BasePostgresSaver):
async with self._cursor(pipeline=True) as cur:
await cur.executemany(query, params)
async def adelete_thread(self, thread_id: str) -> None:
"""Delete all checkpoints and writes associated with a thread ID.
Args:
thread_id (str): The thread ID to delete.
Returns:
None
"""
async with self._cursor(pipeline=True) as cur:
await cur.execute(
"DELETE FROM checkpoints WHERE thread_id = %s",
(str(thread_id),),
)
await cur.execute(
"DELETE FROM checkpoint_blobs WHERE thread_id = %s",
(str(thread_id),),
)
await cur.execute(
"DELETE FROM checkpoint_writes WHERE thread_id = %s",
(str(thread_id),),
)
@asynccontextmanager
async def _cursor(
self, *, pipeline: bool = False
@@ -481,5 +504,30 @@ class AsyncPostgresSaver(BasePostgresSaver):
self.aput_writes(config, writes, task_id, task_path), self.loop
).result()
def delete_thread(self, thread_id: str) -> None:
"""Delete all checkpoints and writes associated with a thread ID.
Args:
thread_id (str): The thread ID to delete.
Returns:
None
"""
try:
# check if we are in the main thread, only bg threads can block
# we don't check in other methods to avoid the overhead
if asyncio.get_running_loop() is self.loop:
raise asyncio.InvalidStateError(
"Synchronous calls to AsyncPostgresSaver are only allowed from a "
"different thread. From the main thread, use the async interface. "
"For example, use `await checkpointer.aget_tuple(...)` or `await "
"graph.ainvoke(...)`."
)
except RuntimeError:
pass
return asyncio.run_coroutine_threadsafe(
self.adelete_thread(thread_id), self.loop
).result()
__all__ = ["AsyncPostgresSaver", "AsyncShallowPostgresSaver", "Conn"]