mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 13:05:15 +02:00
Replace hardcoded database saver class names with `cls` in `from_conn_string` factory methods to improve subclassing support ## Changes * Replaced direct class instantiations with `cls(conn)` in `from_conn_string` classmethods across all database implementations * Updated both synchronous and asynchronous variants for DuckDB, PostgreSQL, and SQLite savers ## Why This refactor makes the database saver classes more extensible by following Python's convention of using `cls` in class methods. This enables proper inheritance patterns where subclasses can reuse the factory methods without needing to override them. Previously, the hardcoded class names would always instantiate the parent class, even when called from a subclass. ## Testing The change is backward compatible and doesn't alter existing functionality. All existing tests should continue to pass as this is purely a structural refactoring that preserves the current behavior while improving extensibility. ## Notes This PR addresses follow up on comments from #2518 - AsyncPostgresSaver didn't need to be fixed but many of the other DB saver classes did.
357 lines
13 KiB
Python
357 lines
13 KiB
Python
import threading
|
|
from contextlib import contextmanager
|
|
from typing import Any, Iterator, Optional, Sequence
|
|
|
|
from langchain_core.runnables import RunnableConfig
|
|
|
|
import duckdb
|
|
from langgraph.checkpoint.base import (
|
|
WRITES_IDX_MAP,
|
|
ChannelVersions,
|
|
Checkpoint,
|
|
CheckpointMetadata,
|
|
CheckpointTuple,
|
|
get_checkpoint_id,
|
|
)
|
|
from langgraph.checkpoint.duckdb.base import BaseDuckDBSaver
|
|
from langgraph.checkpoint.serde.base import SerializerProtocol
|
|
|
|
|
|
class DuckDBSaver(BaseDuckDBSaver):
|
|
lock: threading.Lock
|
|
|
|
def __init__(
|
|
self,
|
|
conn: duckdb.DuckDBPyConnection,
|
|
serde: Optional[SerializerProtocol] = None,
|
|
) -> None:
|
|
super().__init__(serde=serde)
|
|
|
|
self.conn = conn
|
|
self.lock = threading.Lock()
|
|
|
|
@classmethod
|
|
@contextmanager
|
|
def from_conn_string(cls, conn_string: str) -> Iterator["DuckDBSaver"]:
|
|
"""Create a new DuckDBSaver instance from a connection string.
|
|
|
|
Args:
|
|
conn_string (str): The DuckDB connection info string.
|
|
|
|
Returns:
|
|
DuckDBSaver: A new DuckDBSaver instance.
|
|
"""
|
|
with duckdb.connect(conn_string) as conn:
|
|
yield cls(conn)
|
|
|
|
def setup(self) -> None:
|
|
"""Set up the checkpoint database asynchronously.
|
|
|
|
This method creates the necessary tables in the DuckDB database if they don't
|
|
already exist and runs database migrations. It MUST be called directly by the user
|
|
the first time checkpointer is used.
|
|
"""
|
|
with self.lock, self.conn.cursor() as cur:
|
|
try:
|
|
row = cur.execute(
|
|
"SELECT v FROM checkpoint_migrations ORDER BY v DESC LIMIT 1"
|
|
).fetchone()
|
|
if row is None:
|
|
version = -1
|
|
else:
|
|
version = row[0]
|
|
except duckdb.CatalogException:
|
|
version = -1
|
|
for v, migration in zip(
|
|
range(version + 1, len(self.MIGRATIONS)),
|
|
self.MIGRATIONS[version + 1 :],
|
|
):
|
|
cur.execute(migration)
|
|
cur.execute("INSERT INTO checkpoint_migrations (v) VALUES (?)", [v])
|
|
|
|
def list(
|
|
self,
|
|
config: Optional[RunnableConfig],
|
|
*,
|
|
filter: Optional[dict[str, Any]] = None,
|
|
before: Optional[RunnableConfig] = None,
|
|
limit: Optional[int] = None,
|
|
) -> Iterator[CheckpointTuple]:
|
|
"""List checkpoints from the database.
|
|
|
|
This method retrieves a list of checkpoint tuples from the DuckDB database based
|
|
on the provided config. The checkpoints are ordered by checkpoint ID in descending order (newest first).
|
|
|
|
Args:
|
|
config (RunnableConfig): The config to use for listing the checkpoints.
|
|
filter (Optional[Dict[str, Any]]): Additional filtering criteria for metadata. Defaults to None.
|
|
before (Optional[RunnableConfig]): If provided, only checkpoints before the specified checkpoint ID are returned. Defaults to None.
|
|
limit (Optional[int]): The maximum number of checkpoints to return. Defaults to None.
|
|
|
|
Yields:
|
|
Iterator[CheckpointTuple]: An iterator of checkpoint tuples.
|
|
|
|
Examples:
|
|
>>> from langgraph.checkpoint.duckdb import DuckDBSaver
|
|
>>> with DuckDBSaver.from_conn_string(":memory:") as memory:
|
|
... # Run a graph, then list the checkpoints
|
|
>>> config = {"configurable": {"thread_id": "1"}}
|
|
>>> checkpoints = list(memory.list(config, limit=2))
|
|
>>> print(checkpoints)
|
|
[CheckpointTuple(...), CheckpointTuple(...)]
|
|
|
|
>>> config = {"configurable": {"thread_id": "1"}}
|
|
>>> before = {"configurable": {"checkpoint_id": "1ef4f797-8335-6428-8001-8a1503f9b875"}}
|
|
>>> with DuckDBSaver.from_conn_string(":memory:") as memory:
|
|
... # Run a graph, then list the checkpoints
|
|
>>> checkpoints = list(memory.list(config, before=before))
|
|
>>> print(checkpoints)
|
|
[CheckpointTuple(...), ...]
|
|
"""
|
|
where, args = self._search_where(config, filter, before)
|
|
query = self.SELECT_SQL + where + " ORDER BY checkpoint_id DESC"
|
|
if limit:
|
|
query += f" LIMIT {limit}"
|
|
# if we change this to use .stream() we need to make sure to close the cursor
|
|
with self._cursor() as cur:
|
|
cur.execute(query, args)
|
|
for value in cur.fetchall():
|
|
(
|
|
thread_id,
|
|
checkpoint,
|
|
checkpoint_ns,
|
|
checkpoint_id,
|
|
parent_checkpoint_id,
|
|
metadata,
|
|
channel_values,
|
|
pending_writes,
|
|
pending_sends,
|
|
) = value
|
|
yield CheckpointTuple(
|
|
{
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": checkpoint_id,
|
|
}
|
|
},
|
|
self._load_checkpoint(
|
|
checkpoint,
|
|
channel_values,
|
|
pending_sends,
|
|
),
|
|
self._load_metadata(metadata),
|
|
(
|
|
{
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": parent_checkpoint_id,
|
|
}
|
|
}
|
|
if parent_checkpoint_id
|
|
else None
|
|
),
|
|
self._load_writes(pending_writes),
|
|
)
|
|
|
|
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
|
|
"""Get a checkpoint tuple from the database.
|
|
|
|
This method retrieves a checkpoint tuple from the DuckDB database based on the
|
|
provided config. If the config contains a "checkpoint_id" key, the checkpoint with
|
|
the matching thread ID and timestamp is retrieved. Otherwise, the latest checkpoint
|
|
for the given thread ID is retrieved.
|
|
|
|
Args:
|
|
config (RunnableConfig): The config to use for retrieving the checkpoint.
|
|
|
|
Returns:
|
|
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
|
|
|
|
Examples:
|
|
|
|
Basic:
|
|
>>> config = {"configurable": {"thread_id": "1"}}
|
|
>>> checkpoint_tuple = memory.get_tuple(config)
|
|
>>> print(checkpoint_tuple)
|
|
CheckpointTuple(...)
|
|
|
|
With timestamp:
|
|
|
|
>>> config = {
|
|
... "configurable": {
|
|
... "thread_id": "1",
|
|
... "checkpoint_ns": "",
|
|
... "checkpoint_id": "1ef4f797-8335-6428-8001-8a1503f9b875",
|
|
... }
|
|
... }
|
|
>>> checkpoint_tuple = memory.get_tuple(config)
|
|
>>> print(checkpoint_tuple)
|
|
CheckpointTuple(...)
|
|
""" # noqa
|
|
thread_id = config["configurable"]["thread_id"]
|
|
checkpoint_id = get_checkpoint_id(config)
|
|
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
|
|
if checkpoint_id:
|
|
args: tuple[Any, ...] = (thread_id, checkpoint_ns, checkpoint_id)
|
|
where = "WHERE thread_id = ? AND checkpoint_ns = ? AND checkpoint_id = ?"
|
|
else:
|
|
args = (thread_id, checkpoint_ns)
|
|
where = "WHERE thread_id = ? AND checkpoint_ns = ? ORDER BY checkpoint_id DESC LIMIT 1"
|
|
|
|
with self._cursor() as cur:
|
|
cur.execute(
|
|
self.SELECT_SQL + where,
|
|
args,
|
|
)
|
|
|
|
value = cur.fetchone()
|
|
if value:
|
|
(
|
|
thread_id,
|
|
checkpoint,
|
|
checkpoint_ns,
|
|
checkpoint_id,
|
|
parent_checkpoint_id,
|
|
metadata,
|
|
channel_values,
|
|
pending_writes,
|
|
pending_sends,
|
|
) = value
|
|
return CheckpointTuple(
|
|
{
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": checkpoint_id,
|
|
}
|
|
},
|
|
self._load_checkpoint(
|
|
checkpoint,
|
|
channel_values,
|
|
pending_sends,
|
|
),
|
|
self._load_metadata(metadata),
|
|
(
|
|
{
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": parent_checkpoint_id,
|
|
}
|
|
}
|
|
if parent_checkpoint_id
|
|
else None
|
|
),
|
|
self._load_writes(pending_writes),
|
|
)
|
|
|
|
def put(
|
|
self,
|
|
config: RunnableConfig,
|
|
checkpoint: Checkpoint,
|
|
metadata: CheckpointMetadata,
|
|
new_versions: ChannelVersions,
|
|
) -> RunnableConfig:
|
|
"""Save a checkpoint to the database.
|
|
|
|
This method saves a checkpoint to the DuckDB database. The checkpoint is associated
|
|
with the provided config and its parent config (if any).
|
|
|
|
Args:
|
|
config (RunnableConfig): The config to associate with the checkpoint.
|
|
checkpoint (Checkpoint): The checkpoint to save.
|
|
metadata (CheckpointMetadata): Additional metadata to save with the checkpoint.
|
|
new_versions (ChannelVersions): New channel versions as of this write.
|
|
|
|
Returns:
|
|
RunnableConfig: Updated configuration after storing the checkpoint.
|
|
|
|
Examples:
|
|
|
|
>>> from langgraph.checkpoint.duckdb import DuckDBSaver
|
|
>>> with DuckDBSaver.from_conn_string(":memory:") as memory:
|
|
>>> config = {"configurable": {"thread_id": "1", "checkpoint_ns": ""}}
|
|
>>> checkpoint = {"ts": "2024-05-04T06:32:42.235444+00:00", "id": "1ef4f797-8335-6428-8001-8a1503f9b875", "channel_values": {"key": "value"}}
|
|
>>> saved_config = memory.put(config, checkpoint, {"source": "input", "step": 1, "writes": {"key": "value"}}, {})
|
|
>>> print(saved_config)
|
|
{'configurable': {'thread_id': '1', 'checkpoint_ns': '', 'checkpoint_id': '1ef4f797-8335-6428-8001-8a1503f9b875'}}
|
|
"""
|
|
configurable = config["configurable"].copy()
|
|
thread_id = configurable.pop("thread_id")
|
|
checkpoint_ns = configurable.pop("checkpoint_ns")
|
|
checkpoint_id = configurable.pop(
|
|
"checkpoint_id", configurable.pop("thread_ts", None)
|
|
)
|
|
|
|
copy = checkpoint.copy()
|
|
next_config = {
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": checkpoint["id"],
|
|
}
|
|
}
|
|
checkpoint_blobs = self._dump_blobs(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
copy.pop("channel_values"), # type: ignore[misc]
|
|
new_versions,
|
|
)
|
|
with self._cursor() as cur:
|
|
if checkpoint_blobs:
|
|
cur.executemany(self.UPSERT_CHECKPOINT_BLOBS_SQL, checkpoint_blobs)
|
|
cur.execute(
|
|
self.UPSERT_CHECKPOINTS_SQL,
|
|
(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
checkpoint["id"],
|
|
checkpoint_id,
|
|
self._dump_checkpoint(copy),
|
|
self._dump_metadata(metadata),
|
|
),
|
|
)
|
|
return next_config
|
|
|
|
def put_writes(
|
|
self,
|
|
config: RunnableConfig,
|
|
writes: Sequence[tuple[str, Any]],
|
|
task_id: str,
|
|
) -> None:
|
|
"""Store intermediate writes linked to a checkpoint.
|
|
|
|
This method saves intermediate writes associated with a checkpoint to the DuckDB database.
|
|
|
|
Args:
|
|
config (RunnableConfig): Configuration of the related checkpoint.
|
|
writes (List[Tuple[str, Any]]): List of writes to store.
|
|
task_id (str): Identifier for the task creating the writes.
|
|
"""
|
|
query = (
|
|
self.UPSERT_CHECKPOINT_WRITES_SQL
|
|
if all(w[0] in WRITES_IDX_MAP for w in writes)
|
|
else self.INSERT_CHECKPOINT_WRITES_SQL
|
|
)
|
|
with self._cursor() as cur:
|
|
cur.executemany(
|
|
query,
|
|
self._dump_writes(
|
|
config["configurable"]["thread_id"],
|
|
config["configurable"]["checkpoint_ns"],
|
|
config["configurable"]["checkpoint_id"],
|
|
task_id,
|
|
writes,
|
|
),
|
|
)
|
|
|
|
@contextmanager
|
|
def _cursor(self) -> Iterator[duckdb.DuckDBPyConnection]:
|
|
with self.lock, self.conn.cursor() as cur:
|
|
yield cur
|
|
|
|
|
|
__all__ = ["DuckDBSaver", "Conn"]
|