mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-20 22:52:29 +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)
```
291 lines
9.7 KiB
Python
291 lines
9.7 KiB
Python
import random
|
|
from collections.abc import Sequence
|
|
from typing import Any, Optional, cast
|
|
|
|
from langchain_core.runnables import RunnableConfig
|
|
from psycopg.types.json import Jsonb
|
|
|
|
from langgraph.checkpoint.base import (
|
|
WRITES_IDX_MAP,
|
|
BaseCheckpointSaver,
|
|
ChannelVersions,
|
|
Checkpoint,
|
|
CheckpointMetadata,
|
|
get_checkpoint_id,
|
|
)
|
|
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
|
|
from langgraph.checkpoint.serde.types import TASKS, ChannelProtocol
|
|
|
|
MetadataInput = Optional[dict[str, Any]]
|
|
|
|
"""
|
|
To add a new migration, add a new string to the MIGRATIONS list.
|
|
The position of the migration in the list is the version number.
|
|
"""
|
|
MIGRATIONS = [
|
|
"""CREATE TABLE IF NOT EXISTS checkpoint_migrations (
|
|
v INTEGER PRIMARY KEY
|
|
);""",
|
|
"""CREATE TABLE IF NOT EXISTS checkpoints (
|
|
thread_id TEXT NOT NULL,
|
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
checkpoint_id TEXT NOT NULL,
|
|
parent_checkpoint_id TEXT,
|
|
type TEXT,
|
|
checkpoint JSONB NOT NULL,
|
|
metadata JSONB NOT NULL DEFAULT '{}',
|
|
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id)
|
|
);""",
|
|
"""CREATE TABLE IF NOT EXISTS checkpoint_blobs (
|
|
thread_id TEXT NOT NULL,
|
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
channel TEXT NOT NULL,
|
|
version TEXT NOT NULL,
|
|
type TEXT NOT NULL,
|
|
blob BYTEA,
|
|
PRIMARY KEY (thread_id, checkpoint_ns, channel, version)
|
|
);""",
|
|
"""CREATE TABLE IF NOT EXISTS checkpoint_writes (
|
|
thread_id TEXT NOT NULL,
|
|
checkpoint_ns TEXT NOT NULL DEFAULT '',
|
|
checkpoint_id TEXT NOT NULL,
|
|
task_id TEXT NOT NULL,
|
|
idx INTEGER NOT NULL,
|
|
channel TEXT NOT NULL,
|
|
type TEXT,
|
|
blob BYTEA NOT NULL,
|
|
PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx)
|
|
);""",
|
|
"ALTER TABLE checkpoint_blobs ALTER COLUMN blob DROP not null;",
|
|
]
|
|
|
|
SELECT_SQL = f"""
|
|
select
|
|
thread_id,
|
|
checkpoint,
|
|
checkpoint_ns,
|
|
checkpoint_id,
|
|
parent_checkpoint_id,
|
|
metadata,
|
|
(
|
|
select array_agg(array[bl.channel::bytea, bl.type::bytea, bl.blob])
|
|
from jsonb_each_text(checkpoint -> 'channel_versions')
|
|
inner join checkpoint_blobs bl
|
|
on bl.thread_id = checkpoints.thread_id
|
|
and bl.checkpoint_ns = checkpoints.checkpoint_ns
|
|
and bl.channel = jsonb_each_text.key
|
|
and bl.version = jsonb_each_text.value
|
|
) as channel_values,
|
|
(
|
|
select
|
|
array_agg(array[cw.task_id::text::bytea, cw.channel::bytea, cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
|
from checkpoint_writes cw
|
|
where cw.thread_id = checkpoints.thread_id
|
|
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
|
and cw.checkpoint_id = checkpoints.checkpoint_id
|
|
) as pending_writes,
|
|
(
|
|
select array_agg(array[cw.type::bytea, cw.blob] order by cw.task_id, cw.idx)
|
|
from checkpoint_writes cw
|
|
where cw.thread_id = checkpoints.thread_id
|
|
and cw.checkpoint_ns = checkpoints.checkpoint_ns
|
|
and cw.checkpoint_id = checkpoints.parent_checkpoint_id
|
|
and cw.channel = '{TASKS}'
|
|
) as pending_sends
|
|
from checkpoints """
|
|
|
|
UPSERT_CHECKPOINT_BLOBS_SQL = """
|
|
INSERT INTO checkpoint_blobs (thread_id, checkpoint_ns, channel, version, type, blob)
|
|
VALUES (%s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, channel, version) DO NOTHING
|
|
"""
|
|
|
|
UPSERT_CHECKPOINTS_SQL = """
|
|
INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, parent_checkpoint_id, checkpoint, metadata)
|
|
VALUES (%s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id)
|
|
DO UPDATE SET
|
|
checkpoint = EXCLUDED.checkpoint,
|
|
metadata = EXCLUDED.metadata;
|
|
"""
|
|
|
|
UPSERT_CHECKPOINT_WRITES_SQL = """
|
|
INSERT INTO checkpoint_writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) DO UPDATE SET
|
|
channel = EXCLUDED.channel,
|
|
type = EXCLUDED.type,
|
|
blob = EXCLUDED.blob;
|
|
"""
|
|
|
|
INSERT_CHECKPOINT_WRITES_SQL = """
|
|
INSERT INTO checkpoint_writes (thread_id, checkpoint_ns, checkpoint_id, task_id, idx, channel, type, blob)
|
|
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
|
ON CONFLICT (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) DO NOTHING
|
|
"""
|
|
|
|
|
|
class BasePostgresSaver(BaseCheckpointSaver[str]):
|
|
SELECT_SQL = SELECT_SQL
|
|
MIGRATIONS = MIGRATIONS
|
|
UPSERT_CHECKPOINT_BLOBS_SQL = UPSERT_CHECKPOINT_BLOBS_SQL
|
|
UPSERT_CHECKPOINTS_SQL = UPSERT_CHECKPOINTS_SQL
|
|
UPSERT_CHECKPOINT_WRITES_SQL = UPSERT_CHECKPOINT_WRITES_SQL
|
|
INSERT_CHECKPOINT_WRITES_SQL = INSERT_CHECKPOINT_WRITES_SQL
|
|
|
|
jsonplus_serde = JsonPlusSerializer()
|
|
supports_pipeline: bool
|
|
|
|
def _load_checkpoint(
|
|
self,
|
|
checkpoint: dict[str, Any],
|
|
channel_values: list[tuple[bytes, bytes, bytes]],
|
|
pending_sends: list[tuple[bytes, bytes]],
|
|
) -> Checkpoint:
|
|
return {
|
|
**checkpoint,
|
|
"pending_sends": [
|
|
self.serde.loads_typed((c.decode(), b)) for c, b in pending_sends or []
|
|
],
|
|
"channel_values": self._load_blobs(channel_values),
|
|
}
|
|
|
|
def _dump_checkpoint(self, checkpoint: Checkpoint) -> dict[str, Any]:
|
|
return {**checkpoint, "pending_sends": []}
|
|
|
|
def _load_blobs(
|
|
self, blob_values: list[tuple[bytes, bytes, bytes]]
|
|
) -> dict[str, Any]:
|
|
if not blob_values:
|
|
return {}
|
|
return {
|
|
k.decode(): self.serde.loads_typed((t.decode(), v))
|
|
for k, t, v in blob_values
|
|
if t.decode() != "empty"
|
|
}
|
|
|
|
def _dump_blobs(
|
|
self,
|
|
thread_id: str,
|
|
checkpoint_ns: str,
|
|
values: dict[str, Any],
|
|
versions: ChannelVersions,
|
|
) -> list[tuple[str, str, str, str, str, Optional[bytes]]]:
|
|
if not versions:
|
|
return []
|
|
|
|
return [
|
|
(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
k,
|
|
cast(str, ver),
|
|
*(
|
|
self.serde.dumps_typed(values[k])
|
|
if k in values
|
|
else ("empty", None)
|
|
),
|
|
)
|
|
for k, ver in versions.items()
|
|
]
|
|
|
|
def _load_writes(
|
|
self, writes: list[tuple[bytes, bytes, bytes, bytes]]
|
|
) -> list[tuple[str, str, Any]]:
|
|
return (
|
|
[
|
|
(
|
|
tid.decode(),
|
|
channel.decode(),
|
|
self.serde.loads_typed((t.decode(), v)),
|
|
)
|
|
for tid, channel, t, v in writes
|
|
]
|
|
if writes
|
|
else []
|
|
)
|
|
|
|
def _dump_writes(
|
|
self,
|
|
thread_id: str,
|
|
checkpoint_ns: str,
|
|
checkpoint_id: str,
|
|
task_id: str,
|
|
writes: Sequence[tuple[str, Any]],
|
|
) -> list[tuple[str, str, str, str, int, str, str, bytes]]:
|
|
return [
|
|
(
|
|
thread_id,
|
|
checkpoint_ns,
|
|
checkpoint_id,
|
|
task_id,
|
|
WRITES_IDX_MAP.get(channel, idx),
|
|
channel,
|
|
*self.serde.dumps_typed(value),
|
|
)
|
|
for idx, (channel, value) in enumerate(writes)
|
|
]
|
|
|
|
def _load_metadata(self, metadata: dict[str, Any]) -> CheckpointMetadata:
|
|
return self.jsonplus_serde.loads(self.jsonplus_serde.dumps(metadata))
|
|
|
|
def _dump_metadata(self, metadata: CheckpointMetadata) -> str:
|
|
serialized_metadata = self.jsonplus_serde.dumps(metadata)
|
|
# NOTE: we're using JSON serializer (not msgpack), so we need to remove null characters before writing
|
|
return serialized_metadata.decode().replace("\\u0000", "")
|
|
|
|
def get_next_version(self, current: Optional[str], channel: ChannelProtocol) -> str:
|
|
if current is None:
|
|
current_v = 0
|
|
elif isinstance(current, int):
|
|
current_v = current
|
|
else:
|
|
current_v = int(current.split(".")[0])
|
|
next_v = current_v + 1
|
|
next_h = random.random()
|
|
return f"{next_v:032}.{next_h:016}"
|
|
|
|
def _search_where(
|
|
self,
|
|
config: Optional[RunnableConfig],
|
|
filter: MetadataInput,
|
|
before: Optional[RunnableConfig] = None,
|
|
) -> tuple[str, list[Any]]:
|
|
"""Return WHERE clause predicates for alist() given config, filter, before.
|
|
|
|
This method returns a tuple of a string and a tuple of values. The string
|
|
is the parametered WHERE clause predicate (including the WHERE keyword):
|
|
"WHERE column1 = $1 AND column2 IS $2". The list of values contains the
|
|
values for each of the corresponding parameters.
|
|
"""
|
|
wheres = []
|
|
param_values = []
|
|
|
|
# construct predicate for config filter
|
|
if config:
|
|
wheres.append("thread_id = %s ")
|
|
param_values.append(config["configurable"]["thread_id"])
|
|
checkpoint_ns = config["configurable"].get("checkpoint_ns")
|
|
if checkpoint_ns is not None:
|
|
wheres.append("checkpoint_ns = %s")
|
|
param_values.append(checkpoint_ns)
|
|
|
|
if checkpoint_id := get_checkpoint_id(config):
|
|
wheres.append("checkpoint_id = %s ")
|
|
param_values.append(checkpoint_id)
|
|
|
|
# construct predicate for metadata filter
|
|
if filter:
|
|
wheres.append("metadata @> %s ")
|
|
param_values.append(Jsonb(filter))
|
|
|
|
# construct predicate for `before`
|
|
if before is not None:
|
|
wheres.append("checkpoint_id < %s ")
|
|
param_values.append(get_checkpoint_id(before))
|
|
|
|
return (
|
|
"WHERE " + " AND ".join(wheres) if wheres else "",
|
|
param_values,
|
|
)
|