diff --git a/libs/checkpoint-postgres/README.md b/libs/checkpoint-postgres/README.md index 809466262..560cf4577 100644 --- a/libs/checkpoint-postgres/README.md +++ b/libs/checkpoint-postgres/README.md @@ -4,6 +4,9 @@ Implementation of LangGraph CheckpointSaver that uses Postgres. ## Usage +> [!IMPORTANT] +> When using Postgres checkpointers for the first time, make sure to call `.setup()` method on them to create required tables. See example below. + ```python from langgraph.checkpoint.postgres import PostgresSaver @@ -12,6 +15,8 @@ read_config = {"configurable": {"thread_id": "1"}} DB_URI = "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable" with PostgresSaver.from_conn_string(DB_URI) as checkpointer: + # call .setup() the first time you're using the checkpointer + checkpointer.setup() checkpoint = { "v": 1, "ts": "2024-07-31T20:14:19.804150+00:00", diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py index ee7abe120..e07f0aa5f 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/__init__.py @@ -61,9 +61,9 @@ class PostgresSaver(BasePostgresSaver): def setup(self) -> None: """Set up the checkpoint database asynchronously. - This method creates the necessary tables in the SQLite database if they don't - already exist. It is called automatically when needed and should not be called - directly by the user. + This method creates the necessary tables in the Postgres 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: with self.conn.cursor(binary=True) as cur: @@ -90,6 +90,38 @@ class PostgresSaver(BasePostgresSaver): 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 Postgres 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.postgres import PostgresSaver + >>> DB_URI = "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable" + >>> with PostgresSaver.from_conn_string(DB_URI) 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 PostgresSaver.from_conn_string(DB_URI) 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: @@ -121,6 +153,40 @@ class PostgresSaver(BasePostgresSaver): ) def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + """Get a checkpoint tuple from the database. + + This method retrieves a checkpoint tuple from the Postgres 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", "") @@ -171,6 +237,31 @@ class PostgresSaver(BasePostgresSaver): metadata: CheckpointMetadata, new_versions: ChannelVersions, ) -> RunnableConfig: + """Save a checkpoint to the database. + + This method saves a checkpoint to the Postgres 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.postgres import PostgresSaver + >>> DB_URI = "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable" + >>> with PostgresSaver.from_conn_string(DB_URI) 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", "data": {"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") @@ -216,6 +307,15 @@ class PostgresSaver(BasePostgresSaver): writes: List[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 Postgres 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. + """ with self._cursor() as cur: cur.executemany( self.UPSERT_CHECKPOINT_WRITES_SQL, diff --git a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py index 3622150e3..e8b65322b 100644 --- a/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/checkpoint/postgres/aio.py @@ -59,9 +59,9 @@ class AsyncPostgresSaver(BasePostgresSaver): async def setup(self) -> None: """Set up the checkpoint database asynchronously. - This method creates the necessary tables in the SQLite database if they don't - already exist. It is called automatically when needed and should not be called - directly by the user. + This method creates the necessary tables in the Postgres 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. """ async with self.lock: async with self.conn.cursor(binary=True) as cur: @@ -92,6 +92,20 @@ class AsyncPostgresSaver(BasePostgresSaver): before: Optional[RunnableConfig] = None, limit: Optional[int] = None, ) -> AsyncIterator[CheckpointTuple]: + """List checkpoints from the database asynchronously. + + This method retrieves a list of checkpoint tuples from the Postgres database based + on the provided config. The checkpoints are ordered by checkpoint ID in descending order (newest first). + + Args: + config (Optional[RunnableConfig]): Base configuration for filtering checkpoints. + filter (Optional[Dict[str, Any]]): Additional filtering criteria for metadata. + before (Optional[RunnableConfig]): If provided, only checkpoints before the specified checkpoint ID are returned. Defaults to None. + limit (Optional[int]): Maximum number of checkpoints to return. + + Yields: + AsyncIterator[CheckpointTuple]: An asynchronous iterator of matching checkpoint tuples. + """ where, args = self._search_where(config, filter, before) query = self.SELECT_SQL + where + " ORDER BY checkpoint_id DESC" if limit: @@ -125,6 +139,19 @@ class AsyncPostgresSaver(BasePostgresSaver): ) async def aget_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: + """Get a checkpoint tuple from the database asynchronously. + + This method retrieves a checkpoint tuple from the Postgres database based on the + provided config. If the config contains a "checkpoint_id" key, the checkpoint with + the matching thread ID and "checkpoint_id" 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. + """ thread_id = config["configurable"]["thread_id"] checkpoint_id = get_checkpoint_id(config) checkpoint_ns = config["configurable"].get("checkpoint_ns", "") @@ -177,6 +204,20 @@ class AsyncPostgresSaver(BasePostgresSaver): metadata: CheckpointMetadata, new_versions: ChannelVersions, ) -> RunnableConfig: + """Save a checkpoint to the database asynchronously. + + This method saves a checkpoint to the Postgres 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. + """ configurable = config["configurable"].copy() thread_id = configurable.pop("thread_id") checkpoint_ns = configurable.pop("checkpoint_ns") @@ -223,6 +264,15 @@ class AsyncPostgresSaver(BasePostgresSaver): writes: list[tuple[str, Any]], task_id: str, ) -> None: + """Store intermediate writes linked to a checkpoint asynchronously. + + This method saves intermediate writes associated with a checkpoint to the database. + + Args: + config (RunnableConfig): Configuration of the related checkpoint. + writes (Sequence[Tuple[str, Any]]): List of writes to store, each as (channel, value) pair. + task_id (str): Identifier for the task creating the writes. + """ async with self._cursor() as cur: await cur.executemany( self.UPSERT_CHECKPOINT_WRITES_SQL, diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py index d23fa888e..ac2f64f32 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/__init__.py @@ -176,7 +176,7 @@ class SqliteSaver(BaseCheckpointSaver): This method retrieves a checkpoint tuple from the SQLite 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 + the matching thread ID and checkpoint ID is retrieved. Otherwise, the latest checkpoint for the given thread ID is retrieved. Args: @@ -193,12 +193,12 @@ class SqliteSaver(BaseCheckpointSaver): >>> print(checkpoint_tuple) CheckpointTuple(...) - With timestamp: + With checkpoint ID: >>> config = { ... "configurable": { ... "thread_id": "1", - ... "checkpoint_ns": "", + ... "checkpoint_ns": "", ... "checkpoint_id": "1ef4f797-8335-6428-8001-8a1503f9b875", ... } ... } @@ -283,12 +283,12 @@ class SqliteSaver(BaseCheckpointSaver): """List checkpoints from the database. This method retrieves a list of checkpoint tuples from the SQLite database based - on the provided config. The checkpoints are ordered by timestamp in descending order. + 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 timestamp are returned. 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: @@ -367,10 +367,11 @@ class SqliteSaver(BaseCheckpointSaver): Args: config (RunnableConfig): The config to associate with the checkpoint. checkpoint (Checkpoint): The checkpoint to save. - metadata (Optional[dict[str, Any]]): Additional metadata to save with the checkpoint. Defaults to None. + metadata (CheckpointMetadata): Additional metadata to save with the checkpoint. + new_versions (ChannelVersions): New channel versions as of this write. Returns: - RunnableConfig: The updated config containing the saved checkpoint's timestamp. + RunnableConfig: Updated configuration after storing the checkpoint. Examples: diff --git a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py index f08fe62de..95b77209a 100644 --- a/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py +++ b/libs/checkpoint-sqlite/langgraph/checkpoint/sqlite/aio.py @@ -230,7 +230,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver): This method retrieves a checkpoint tuple from the SQLite 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 + the matching thread ID and checkpoint ID is retrieved. Otherwise, the latest checkpoint for the given thread ID is retrieved. Args: @@ -317,12 +317,12 @@ class AsyncSqliteSaver(BaseCheckpointSaver): """List checkpoints from the database asynchronously. This method retrieves a list of checkpoint tuples from the SQLite database based - on the provided config. The checkpoints are ordered by timestamp in descending order. + on the provided config. The checkpoints are ordered by checkpoint ID in descending order (newest first). Args: config (Optional[RunnableConfig]): Base configuration for filtering checkpoints. filter (Optional[Dict[str, Any]]): Additional filtering criteria for metadata. - before (Optional[RunnableConfig]): List checkpoints created before this configuration. + before (Optional[RunnableConfig]): If provided, only checkpoints before the specified checkpoint ID are returned. Defaults to None. limit (Optional[int]): Maximum number of checkpoints to return. Yields: @@ -385,10 +385,10 @@ class AsyncSqliteSaver(BaseCheckpointSaver): 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 (dict): New versions as of this write + new_versions (ChannelVersions): New channel versions as of this write. Returns: - RunnableConfig: The updated config containing the saved checkpoint's timestamp. + RunnableConfig: Updated configuration after storing the checkpoint. """ await self.setup() thread_id = config["configurable"]["thread_id"] diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index 4f9c94cba..86c8b0eec 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -289,7 +289,7 @@ class BaseCheckpointSaver(ABC): config (RunnableConfig): Configuration for the checkpoint. checkpoint (Checkpoint): The checkpoint to store. metadata (CheckpointMetadata): Additional metadata for the checkpoint. - new_versions (dict): New versions as of this write + new_versions (ChannelVersions): New channel versions as of this write. Returns: RunnableConfig: Updated configuration after storing the checkpoint. @@ -381,7 +381,7 @@ class BaseCheckpointSaver(ABC): config (RunnableConfig): Configuration for the checkpoint. checkpoint (Checkpoint): The checkpoint to store. metadata (CheckpointMetadata): Additional metadata for the checkpoint. - new_versions (dict): New versions as of this write + new_versions (ChannelVersions): New channel versions as of this write. Returns: RunnableConfig: Updated configuration after storing the checkpoint.