mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-24 16:42:24 +02:00
- Initializing the store with an 'embedding config' -> this contains the
'dims' (used to create the table) and the encoder object (rn langchain
embeddings object, though that is ......)
- Call setup() -> creates the vector table.
Each document has 1 or more vectors associated with it for each json
path in the embedding config.
Would welcome critique and requests!
Leaving the params as the defaults for pgvector but open to feedback if
you think it's important to be able to more transparently configure that
in setup()
```python
from typing import TypedDict, List, Dict, Any, Optional
from langchain_openai import OpenAIEmbeddings
from langgraph.graph import StateGraph
from langgraph.store.postgres import PostgresStore
emb_config = {
"dims": 1536, # OpenAI embedding dimensions
"embed": OpenAIEmbeddings(model="text-embedding-3-small"),
"distance_type": "cosine",
}
with PostgresStore.from_conn_string(
"postgres://postgres:postgres@localhost:5441",
embedding=emb_config,
) as store:
store.setup()
# Define the state type for our graph
class State(TypedDict):
query: str
results: Optional[List[Dict[str, Any]]]
def put_stuff(state: State) -> State:
docs = [
("doc1", {"text": "red apple in kitchen"}),
("doc2", {"text": "blue car in garage"}),
("doc3", {"text": "green apple on table"}),
]
for key, value in docs:
store.put(("docs",), key, value)
def search_stuff(state: State) -> State:
"""Search for documents using vector similarity."""
results = store.search(("docs",), query=state["query"])
return {"results": results}
builder = StateGraph(State)
builder.add_node(put_stuff)
builder.add_node(search_stuff)
builder.add_edge("__start__", "put_stuff")
builder.add_edge("put_stuff", "search_stuff")
# Compile
with PostgresStore.from_conn_string(
"postgres://postgres:postgres@localhost:5441",
embedding=emb_config,
) as store:
chain = builder.compile(store=store)
result = chain.invoke({"query": "sour apple"})
# Print results
for doc in result["results"]:
print(doc.key)
print(doc.value)
print(doc.response_metadata)
```
402 lines
16 KiB
Python
402 lines
16 KiB
Python
import threading
|
|
from collections.abc import Iterator, Sequence
|
|
from contextlib import contextmanager
|
|
from typing import Any, Optional
|
|
|
|
from langchain_core.runnables import RunnableConfig
|
|
from psycopg import Capabilities, Connection, Cursor, Pipeline
|
|
from psycopg.errors import UndefinedTable
|
|
from psycopg.rows import DictRow, dict_row
|
|
from psycopg.types.json import Jsonb
|
|
from psycopg_pool import ConnectionPool
|
|
|
|
from langgraph.checkpoint.base import (
|
|
WRITES_IDX_MAP,
|
|
ChannelVersions,
|
|
Checkpoint,
|
|
CheckpointMetadata,
|
|
CheckpointTuple,
|
|
get_checkpoint_id,
|
|
)
|
|
from langgraph.checkpoint.postgres import _internal
|
|
from langgraph.checkpoint.postgres.base import BasePostgresSaver
|
|
from langgraph.checkpoint.serde.base import SerializerProtocol
|
|
|
|
Conn = _internal.Conn # For backward compatibility
|
|
|
|
|
|
class PostgresSaver(BasePostgresSaver):
|
|
lock: threading.Lock
|
|
|
|
def __init__(
|
|
self,
|
|
conn: _internal.Conn,
|
|
pipe: Optional[Pipeline] = None,
|
|
serde: Optional[SerializerProtocol] = None,
|
|
) -> None:
|
|
super().__init__(serde=serde)
|
|
if isinstance(conn, ConnectionPool) and pipe is not None:
|
|
raise ValueError(
|
|
"Pipeline should be used only with a single Connection, not ConnectionPool."
|
|
)
|
|
|
|
self.conn = conn
|
|
self.pipe = pipe
|
|
self.lock = threading.Lock()
|
|
self.supports_pipeline = Capabilities().has_pipeline()
|
|
|
|
@classmethod
|
|
@contextmanager
|
|
def from_conn_string(
|
|
cls, conn_string: str, *, pipeline: bool = False
|
|
) -> Iterator["PostgresSaver"]:
|
|
"""Create a new PostgresSaver instance from a connection string.
|
|
|
|
Args:
|
|
conn_string (str): The Postgres connection info string.
|
|
pipeline (bool): whether to use Pipeline
|
|
|
|
Returns:
|
|
PostgresSaver: A new PostgresSaver instance.
|
|
"""
|
|
with Connection.connect(
|
|
conn_string, autocommit=True, prepare_threshold=0, row_factory=dict_row
|
|
) as conn:
|
|
if pipeline:
|
|
with conn.pipeline() as pipe:
|
|
yield cls(conn, pipe)
|
|
else:
|
|
yield cls(conn)
|
|
|
|
def setup(self) -> None:
|
|
"""Set up the checkpoint database asynchronously.
|
|
|
|
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._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["v"]
|
|
except UndefinedTable:
|
|
version = -1
|
|
for v, migration in zip(
|
|
range(version + 1, len(self.MIGRATIONS)),
|
|
self.MIGRATIONS[version + 1 :],
|
|
):
|
|
cur.execute(migration)
|
|
cur.execute(f"INSERT INTO checkpoint_migrations (v) VALUES ({v})")
|
|
if self.pipe:
|
|
self.pipe.sync()
|
|
|
|
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 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:
|
|
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, binary=True)
|
|
for value in cur:
|
|
yield CheckpointTuple(
|
|
{
|
|
"configurable": {
|
|
"thread_id": value["thread_id"],
|
|
"checkpoint_ns": value["checkpoint_ns"],
|
|
"checkpoint_id": value["checkpoint_id"],
|
|
}
|
|
},
|
|
self._load_checkpoint(
|
|
value["checkpoint"],
|
|
value["channel_values"],
|
|
value["pending_sends"],
|
|
),
|
|
self._load_metadata(value["metadata"]),
|
|
(
|
|
{
|
|
"configurable": {
|
|
"thread_id": value["thread_id"],
|
|
"checkpoint_ns": value["checkpoint_ns"],
|
|
"checkpoint_id": value["parent_checkpoint_id"],
|
|
}
|
|
}
|
|
if value["parent_checkpoint_id"]
|
|
else None
|
|
),
|
|
self._load_writes(value["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 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", "")
|
|
if checkpoint_id:
|
|
args: tuple[Any, ...] = (thread_id, checkpoint_ns, checkpoint_id)
|
|
where = "WHERE thread_id = %s AND checkpoint_ns = %s AND checkpoint_id = %s"
|
|
else:
|
|
args = (thread_id, checkpoint_ns)
|
|
where = "WHERE thread_id = %s AND checkpoint_ns = %s ORDER BY checkpoint_id DESC LIMIT 1"
|
|
|
|
with self._cursor() as cur:
|
|
cur.execute(
|
|
self.SELECT_SQL + where,
|
|
args,
|
|
binary=True,
|
|
)
|
|
|
|
for value in cur:
|
|
return CheckpointTuple(
|
|
{
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": value["checkpoint_id"],
|
|
}
|
|
},
|
|
self._load_checkpoint(
|
|
value["checkpoint"],
|
|
value["channel_values"],
|
|
value["pending_sends"],
|
|
),
|
|
self._load_metadata(value["metadata"]),
|
|
(
|
|
{
|
|
"configurable": {
|
|
"thread_id": thread_id,
|
|
"checkpoint_ns": checkpoint_ns,
|
|
"checkpoint_id": value["parent_checkpoint_id"],
|
|
}
|
|
}
|
|
if value["parent_checkpoint_id"]
|
|
else None
|
|
),
|
|
self._load_writes(value["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 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", "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"],
|
|
}
|
|
}
|
|
|
|
with self._cursor(pipeline=True) as cur:
|
|
cur.executemany(
|
|
self.UPSERT_CHECKPOINT_BLOBS_SQL,
|
|
self._dump_blobs(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
copy.pop("channel_values"), # type: ignore[misc]
|
|
new_versions,
|
|
),
|
|
)
|
|
cur.execute(
|
|
self.UPSERT_CHECKPOINTS_SQL,
|
|
(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
checkpoint["id"],
|
|
checkpoint_id,
|
|
Jsonb(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 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.
|
|
"""
|
|
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(pipeline=True) 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, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]:
|
|
"""Create a database cursor as a context manager.
|
|
|
|
Args:
|
|
pipeline (bool): whether to use pipeline for the DB operations inside the context manager.
|
|
Will be applied regardless of whether the PostgresSaver instance was initialized with a pipeline.
|
|
If pipeline mode is not supported, will fall back to using transaction context manager.
|
|
"""
|
|
with _internal.get_connection(self.conn) as conn:
|
|
if self.pipe:
|
|
# a connection in pipeline mode can be used concurrently
|
|
# in multiple threads/coroutines, but only one cursor can be
|
|
# used at a time
|
|
try:
|
|
with conn.cursor(binary=True, row_factory=dict_row) as cur:
|
|
yield cur
|
|
finally:
|
|
if pipeline:
|
|
self.pipe.sync()
|
|
elif pipeline:
|
|
# a connection not in pipeline mode can only be used by one
|
|
# thread/coroutine at a time, so we acquire a lock
|
|
if self.supports_pipeline:
|
|
with (
|
|
self.lock,
|
|
conn.pipeline(),
|
|
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
|
):
|
|
yield cur
|
|
else:
|
|
# Use connection's transaction context manager when pipeline mode not supported
|
|
with (
|
|
self.lock,
|
|
conn.transaction(),
|
|
conn.cursor(binary=True, row_factory=dict_row) as cur,
|
|
):
|
|
yield cur
|
|
else:
|
|
with self.lock, conn.cursor(binary=True, row_factory=dict_row) as cur:
|
|
yield cur
|
|
|
|
|
|
__all__ = ["PostgresSaver", "BasePostgresSaver", "Conn"]
|