mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 04:55:09 +02:00
- Add Pregel.get_state_history and .aget_state_history methods to get history iterator - Update checkpointer base class with new list and get_tuple methods - Rewrite in-memory checkpointer class to track history - Rewrite sqlite and aiosqlite checkpointers to track history - Add new tests for history tracking
114 lines
4.0 KiB
Python
114 lines
4.0 KiB
Python
import pickle
|
|
from contextlib import AbstractAsyncContextManager
|
|
from types import TracebackType
|
|
from typing import AsyncIterator, Optional, Self
|
|
|
|
import aiosqlite
|
|
from langchain_core.pydantic_v1 import Field
|
|
from langchain_core.runnables import RunnableConfig
|
|
|
|
from langgraph.checkpoint.base import BaseCheckpointSaver, Checkpoint, CheckpointTuple
|
|
|
|
|
|
class AsyncSqliteSaver(BaseCheckpointSaver, AbstractAsyncContextManager):
|
|
conn: aiosqlite.Connection
|
|
|
|
is_setup: bool = Field(False, init=False, repr=False)
|
|
|
|
class Config:
|
|
arbitrary_types_allowed = True
|
|
|
|
@classmethod
|
|
def from_conn_string(cls, conn_string: str) -> "AsyncSqliteSaver":
|
|
return AsyncSqliteSaver(conn=aiosqlite.connect(conn_string))
|
|
|
|
async def __aenter__(self) -> Self:
|
|
return self
|
|
|
|
async def __aexit__(
|
|
self,
|
|
__exc_type: type[BaseException] | None,
|
|
__exc_value: BaseException | None,
|
|
__traceback: TracebackType | None,
|
|
) -> bool | None:
|
|
return await self.conn.close()
|
|
|
|
async def setup(self) -> None:
|
|
if self.is_setup:
|
|
return
|
|
|
|
await self.conn
|
|
async with self.conn.executescript(
|
|
"""
|
|
CREATE TABLE IF NOT EXISTS checkpoints (
|
|
thread_id TEXT NOT NULL,
|
|
thread_ts TEXT NOT NULL,
|
|
checkpoint BLOB,
|
|
PRIMARY KEY (thread_id, thread_ts)
|
|
);
|
|
"""
|
|
):
|
|
await self.conn.commit()
|
|
|
|
self.is_setup = True
|
|
|
|
async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
|
await self.setup()
|
|
if config["configurable"].get("thread_ts"):
|
|
async with self.conn.execute(
|
|
"SELECT checkpoint 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]))
|
|
else:
|
|
async with self.conn.execute(
|
|
"SELECT thread_id, thread_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():
|
|
return CheckpointTuple(
|
|
{
|
|
"configurable": {
|
|
"thread_id": value[0],
|
|
"thread_ts": value[1],
|
|
}
|
|
},
|
|
pickle.loads(value[2]),
|
|
)
|
|
|
|
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",
|
|
(config["configurable"]["thread_id"],),
|
|
) as cursor:
|
|
async for thread_id, thread_ts, value in cursor:
|
|
yield CheckpointTuple(
|
|
{"configurable": {"thread_id": thread_id, "thread_ts": thread_ts}},
|
|
pickle.loads(value),
|
|
)
|
|
|
|
async def aput(
|
|
self, config: RunnableConfig, checkpoint: Checkpoint
|
|
) -> RunnableConfig:
|
|
await self.setup()
|
|
async with self.conn.execute(
|
|
"INSERT OR REPLACE INTO checkpoints (thread_id, thread_ts, checkpoint) VALUES (?, ?, ?)",
|
|
(
|
|
config["configurable"]["thread_id"],
|
|
checkpoint["ts"],
|
|
pickle.dumps(checkpoint),
|
|
),
|
|
):
|
|
await self.conn.commit()
|
|
return {
|
|
"configurable": {
|
|
"thread_id": config["configurable"]["thread_id"],
|
|
"thread_ts": checkpoint["ts"],
|
|
}
|
|
}
|