diff --git a/libs/checkpoint-postgres/Makefile b/libs/checkpoint-postgres/Makefile index 33ed2a3c4..adf92f262 100644 --- a/libs/checkpoint-postgres/Makefile +++ b/libs/checkpoint-postgres/Makefile @@ -5,7 +5,11 @@ ###################### start-postgres: - POSTGRES_VERSION=${POSTGRES_VERSION:-16} docker compose -f tests/compose-postgres.yml up -V --force-recreate --wait + POSTGRES_VERSION=${POSTGRES_VERSION:-16} docker compose -f tests/compose-postgres.yml up -V --force-recreate --wait || ( \ + echo "Failed to start PostgreSQL, printing logs..."; \ + docker compose -f tests/compose-postgres.yml logs; \ + exit 1 \ + ) stop-postgres: docker compose -f tests/compose-postgres.yml down diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py index 28162021c..0f1bda363 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/aio.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/aio.py @@ -19,21 +19,37 @@ from psycopg.rows import DictRow, dict_row from psycopg_pool import AsyncConnectionPool from langgraph.checkpoint.postgres import _ainternal -from langgraph.store.base import GetOp, ListNamespacesOp, Op, PutOp, Result, SearchOp +from langgraph.store.base import ( + GetOp, + ListNamespacesOp, + Op, + PutOp, + Result, + SearchItem, + SearchOp, +) from langgraph.store.base.batch import AsyncBatchedBaseStore from langgraph.store.postgres.base import ( BasePostgresStore, + EmbeddingConfig, Row, _decode_ns_bytes, _group_ops, _row_to_item, + check_vector_available, ) logger = logging.getLogger(__name__) class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Conn]): - __slots__ = ("_deserializer", "pipe", "lock", "supports_pipeline") + __slots__ = ( + "_deserializer", + "pipe", + "lock", + "supports_pipeline", + "embedding_config", + ) def __init__( self, @@ -43,6 +59,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con deserializer: Optional[ Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]] ] = None, + embedding: Optional[EmbeddingConfig] = None, ) -> None: if isinstance(conn, AsyncConnectionPool) and pipe is not None: raise ValueError( @@ -55,6 +72,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con self.lock = asyncio.Lock() self.loop = asyncio.get_running_loop() self.supports_pipeline = Capabilities().has_pipeline() + self.embedding_config = embedding async def abatch(self, ops: Iterable[Op]) -> list[Result]: grouped_ops, num_ops = _group_ops(ops) @@ -75,7 +93,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con results: list[Result], conn: AsyncConnection[DictRow], ) -> None: - async with self._cursor(conn, pipeline=True) as cur: + async with self._cursor(pipeline=True) as cur: if GetOp in grouped_ops: await self._batch_get_ops( cast(Sequence[tuple[int, GetOp]], grouped_ops[GetOp]), @@ -130,7 +148,28 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con put_ops: Sequence[tuple[int, PutOp]], cur: AsyncCursor[DictRow], ) -> None: - queries = self._get_batch_PUT_queries(put_ops) + queries, embedding_request = self._prepare_batch_PUT_queries(put_ops) + if embedding_request: + if self.embedding_config is None: + # Should not get here since the embedding config is required + # to return an embedding_request above + raise ValueError( + "Embedding configuration is required for vector operations " + f"(for semantic search). " + f"Please provide an EmbeddingConfig when initializing the {self.__class__.__name__}." + ) + query, txt_params = embedding_request + # Update the params to replace the raw text with the vectors + vectors = await self.embedding_config["embed"].aembed_documents( + [param[-1] for param in txt_params] + ) + queries.extend( + [ + (query, (ns, key, value, vector)) + for (ns, key, value, _), vector in zip(txt_params, vectors) + ] + ) + for query, params in queries: await cur.execute(query, params) @@ -140,13 +179,24 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con results: list[Result], cur: AsyncCursor[DictRow], ) -> None: - queries = self._get_batch_search_queries(search_ops) - for (query, params), (idx, _) in zip(queries, search_ops): + queries, embedding_requests = self._prepare_batch_search_queries(search_ops) + + if embedding_requests and self.embedding_config: + embeddings = await self.embedding_config["embed"].aembed_documents( + [query for _, query in embedding_requests] + ) + for (idx, _), embedding in zip(embedding_requests, embeddings): + queries[idx][1][0] = embedding + + for (idx, _), (query, params) in zip(search_ops, queries): await cur.execute(query, params) rows = cast(list[Row], await cur.fetchall()) items = [ _row_to_item( - _decode_ns_bytes(row["prefix"]), row, loader=self._deserializer + _decode_ns_bytes(row["prefix"]), + row, + loader=self._deserializer, + cls=SearchItem, ) for row in rows ] @@ -167,40 +217,42 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con @asynccontextmanager async def _cursor( - self, conn: AsyncConnection[DictRow], *, pipeline: bool = False - ) -> AsyncIterator[AsyncCursor[Any]]: + self, *, pipeline: bool = False + ) -> AsyncIterator[AsyncCursor[DictRow]]: """Create a database cursor as a context manager. Args: - conn: The database connection to use pipeline: whether to use pipeline for the DB operations inside the context manager. Will be applied regardless of whether the PostgresStore instance was initialized with a pipeline. If pipeline mode is not supported, will fall back to using transaction context manager. """ - 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 - async with conn.cursor(binary=True) as cur: + async with _ainternal.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: - yield cur + async with conn.cursor(binary=True, row_factory=dict_row) as cur: + yield cur finally: if pipeline: await 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: - async with self.lock, conn.pipeline(), conn.cursor(binary=True) as cur: - yield cur + 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: + async with self.lock, conn.pipeline(), conn.cursor( + binary=True, row_factory=dict_row + ) as cur: + yield cur + else: + async with self.lock, conn.transaction(), conn.cursor( + binary=True, row_factory=dict_row + ) as cur: + yield cur else: - async with self.lock, conn.transaction(), conn.cursor( - binary=True - ) as cur: + async with conn.cursor(binary=True, row_factory=dict_row) as cur: yield cur - else: - async with conn.cursor(binary=True) as cur: - yield cur def batch(self, ops: Iterable[Op]) -> list[Result]: return asyncio.run_coroutine_threadsafe(self.abatch(ops), self.loop).result() @@ -215,6 +267,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con min_size: int = 1, max_size: Optional[int] = None, use_pool: bool = False, + embedding: Optional[EmbeddingConfig] = None, ) -> AsyncIterator["AsyncPostgresStore"]: """Create a new AsyncPostgresStore instance from a connection string. @@ -224,6 +277,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con min_size (int): Minimum number of connections when using a pool max_size (Optional[int]): Maximum number of connections when using a pool use_pool (bool): Whether to use a connection pool + embedding (Optional[EmbeddingConfig]): Configuration for vector embeddings Returns: AsyncPostgresStore: A new AsyncPostgresStore instance. @@ -242,16 +296,16 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con }, ), ) as pool: - yield cls(conn=pool) + yield cls(conn=pool, embedding=embedding) else: async with await AsyncConnection.connect( conn_string, autocommit=True, prepare_threshold=0, row_factory=dict_row ) as conn: if pipeline: async with conn.pipeline() as pipe: - yield cls(conn=conn, pipe=pipe) + yield cls(conn=conn, pipe=pipe, embedding=embedding) else: - yield cls(conn=conn) + yield cls(conn=conn, embedding=embedding) async def setup(self) -> None: """Set up the store database asynchronously. @@ -260,33 +314,59 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con already exist and runs database migrations. It MUST be called directly by the user the first time the store is used. """ - async with _ainternal.get_connection(self.conn) as conn: - async with conn.cursor() as cur: - try: - await cur.execute( - "SELECT v FROM store_migrations ORDER BY v DESC LIMIT 1" - ) - row = cast(dict, await cur.fetchone()) - if row is None: - version = -1 - else: - version = row["v"] - except UndefinedTable: + async with self._cursor() as cur: + try: + await cur.execute( + "SELECT v FROM store_migrations ORDER BY v DESC LIMIT 1" + ) + row = await cur.fetchone() + if row is None: version = -1 - # Create store_migrations table if it doesn't exist - await cur.execute( - """ - CREATE TABLE IF NOT EXISTS store_migrations ( - v INTEGER PRIMARY KEY - ) - """ + else: + version = row["v"] + except UndefinedTable: + version = -1 + await cur.execute( + """ + CREATE TABLE IF NOT EXISTS store_migrations ( + v INTEGER PRIMARY KEY ) - for v, migration in enumerate( - self.MIGRATIONS[version + 1 :], start=version + 1 - ): - await cur.execute(migration) + """ + ) + + for v, migration in enumerate( + self.MIGRATIONS[version + 1 :], start=version + 1 + ): + if isinstance(migration, str): + sql = migration + else: + if migration.acondition and not (await migration.acondition(self)): + continue + + sql = migration.sql + if migration.params: + params = { + k: v(self) if v is not None and callable(v) else v + for k, v in migration.params.items() + } + try: + sql = sql % params + except Exception as e: + logger.warning(f"Failed to format migration {v}: {e}") + if migration.condition == check_vector_available: + self.embedding_config = None + continue + + try: + await cur.execute(sql) await cur.execute( "INSERT INTO store_migrations (v) VALUES (%s)", (v,) ) - if self.pipe: - await self.pipe.sync() + except Exception as e: + logger.warning(f"Failed to run migration {v}: {e}") + if ( + not isinstance(migration, str) + and migration.condition == check_vector_available + ): + self.embedding_config = None + continue diff --git a/libs/checkpoint-postgres/langgraph/store/postgres/base.py b/libs/checkpoint-postgres/langgraph/store/postgres/base.py index 5f9bc5acf..c9df61c26 100644 --- a/libs/checkpoint-postgres/langgraph/store/postgres/base.py +++ b/libs/checkpoint-postgres/langgraph/store/postgres/base.py @@ -4,21 +4,26 @@ import logging import threading from collections import defaultdict from contextlib import contextmanager +from dataclasses import dataclass from datetime import datetime from typing import ( Any, + Awaitable, Callable, Generic, Iterable, Iterator, + List, Optional, Sequence, + Type, TypeVar, Union, cast, ) import orjson +from langchain_core.embeddings import Embeddings from psycopg import Capabilities, Connection, Cursor, Pipeline from psycopg.errors import UndefinedTable from psycopg.rows import DictRow, dict_row @@ -35,13 +40,91 @@ from langgraph.store.base import ( Op, PutOp, Result, + SearchItem, SearchOp, ) logger = logging.getLogger(__name__) -MIGRATIONS = [ +class EmbeddingConfig(TypedDict, total=False): + """Configuration for vector embeddings in PostgreSQL store.""" + + dims: int + """Number of dimensions in the embedding vectors. + + Common embedding models have the following dimensions: + - OpenAI text-embedding-3-large: 256, 1024, or 3072 + - OpenAI text-embedding-3-small: 512 or 1536 + - OpenAI text-embedding-ada-002: 1536 + - Cohere embed-english-v3.0: 1024 + - Cohere embed-english-light-v3.0: 384 + - Cohere embed-multilingual-v3.0: 1024 + - Cohere embed-multilingual-light-v3.0: 384 + """ + + embed: Embeddings + """Optional function to generate embeddings from text.""" + + text_fields: Optional[list[str]] + """Fields to extract text from for embedding generation. + + Defaults to ["__root__"], which embeds the json object as a whole. + """ + + +@dataclass +class Migration: + """A database migration with optional conditions and parameters.""" + + sql: str + condition: Optional[Callable[[Any], bool]] = None + acondition: Optional[Callable[[Any], Awaitable[bool]]] = None + params: Optional[dict[str, Any]] = None + + +def check_vector_available(store: Any) -> bool: + """Check if vector operations are available in the database.""" + if store.embedding_config is None: + # Need the dims to initialize the table + return False + try: + with store._cursor() as cur: + cur.execute( + """ + SELECT 1 FROM pg_available_extensions WHERE name = 'vector' + """ + ) + result = bool(cur.fetchone()) + if not result: + logger.warning("Vector extension is not available in the database.") + return result + except Exception as e: + logger.warning(f"Failed to check vector extension availability: {e}") + return False + + +async def acheck_vector_available(store: Any) -> bool: + if store.embedding_config is None: + # Need the dims to initialize the table + return False + try: + async with store._cursor() as cur: + await cur.execute( + """ + SELECT 1 FROM pg_available_extensions WHERE name = 'vector' + """ + ) + result = bool(await cur.fetchone()) + if not result: + logger.warning("Vector extension is not available in the database.") + return result + except Exception as e: + logger.warning(f"Failed to check vector extension availability: {e}") + return False + + +MIGRATIONS: Sequence[Union[str, Migration]] = [ """ CREATE TABLE IF NOT EXISTS store ( -- 'prefix' represents the doc's 'namespace' @@ -57,6 +140,38 @@ CREATE TABLE IF NOT EXISTS store ( -- For faster lookups by prefix CREATE INDEX IF NOT EXISTS store_prefix_idx ON store USING btree (prefix text_pattern_ops); """, + Migration( + """ +CREATE EXTENSION IF NOT EXISTS vector; +""", + condition=check_vector_available, + acondition=acheck_vector_available, + ), + Migration( + """ +CREATE TABLE IF NOT EXISTS store_vectors ( + prefix text NOT NULL, + key text NOT NULL, + field_name text NOT NULL, + embedding vector(%(dims)s), + created_at TIMESTAMP WITH TIME ZONE DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP WITH TIME ZONE DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (prefix, key, field_name), + FOREIGN KEY (prefix, key) REFERENCES store(prefix, key) ON DELETE CASCADE +); +""", + condition=check_vector_available, + acondition=acheck_vector_available, + params={"dims": lambda store: store.embedding_config["dims"]}, + ), + Migration( + """ +CREATE INDEX IF NOT EXISTS store_vectors_embedding_idx ON store_vectors + USING ivfflat (embedding vector_cosine_ops); +""", + condition=check_vector_available, + acondition=acheck_vector_available, + ), ] C = TypeVar("C", bound=Union[_pg_internal.Conn, _ainternal.Conn]) @@ -66,6 +181,7 @@ class BasePostgresStore(Generic[C]): MIGRATIONS = MIGRATIONS conn: C _deserializer: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]] + embedding_config: Optional[EmbeddingConfig] def _get_batch_GET_ops_queries( self, @@ -87,10 +203,13 @@ class BasePostgresStore(Generic[C]): results.append((query, params, namespace, items)) return results - def _get_batch_PUT_queries( + def _prepare_batch_PUT_queries( self, put_ops: Sequence[tuple[int, PutOp]], - ) -> list[tuple[str, Sequence]]: + ) -> tuple[ + list[tuple[str, Sequence]], + Optional[tuple[str, Sequence[tuple[str, str, str, str]]]], + ]: # Last-write wins dedupped_ops: dict[tuple[tuple[str, ...], str], PutOp] = {} for _, op in put_ops: @@ -117,60 +236,127 @@ class BasePostgresStore(Generic[C]): ) params = (_namespace_to_text(namespace), *keys) queries.append((query, params)) + embedding_request: Optional[tuple[str, Sequence[tuple[str, str, str, str]]]] = ( + None + ) if inserts: values = [] insertion_params = [] + vector_values = [] + embedding_request_params = [] + + # First handle main store insertions for op in inserts: values.append("(%s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)") insertion_params.extend( [ _namespace_to_text(op.namespace), op.key, - Jsonb(op.value), + Jsonb(cast(dict, op.value).copy()), ] ) + + # Then handle embeddings if configured + if self.embedding_config: + text_fields = self.embedding_config.get("text_fields", ["__root__"]) + if isinstance(text_fields, str): + text_fields = [text_fields] + elif text_fields is None: + text_fields = ["__root__"] + + for op in inserts: + value = op.value + ns = _namespace_to_text(op.namespace) + k = op.key + + for field in text_fields: + for text in _extract_text_by_path(value, field): + vector_values.append( + "(%s, %s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)" + ) + embedding_request_params.append((ns, k, field, text)) + values_str = ",".join(values) query = f""" INSERT INTO store (prefix, key, value, created_at, updated_at) VALUES {values_str} ON CONFLICT (prefix, key) DO UPDATE - SET value = EXCLUDED.value, updated_at = CURRENT_TIMESTAMP + SET value = EXCLUDED.value, + updated_at = CURRENT_TIMESTAMP """ queries.append((query, insertion_params)) - return queries + if vector_values: + values_str = ",".join(vector_values) + query = f""" + INSERT INTO store_vectors (prefix, key, field_name, embedding, created_at, updated_at) + VALUES {values_str} + ON CONFLICT (prefix, key, field_name) DO UPDATE + SET embedding = EXCLUDED.embedding, + updated_at = CURRENT_TIMESTAMP + """ + embedding_request = (query, embedding_request_params) - def _get_batch_search_queries( + return queries, embedding_request + + def _prepare_batch_search_queries( self, search_ops: Sequence[tuple[int, SearchOp]], - ) -> list[tuple[str, Sequence]]: - queries: list[tuple[str, Sequence]] = [] - for _, op in search_ops: - query = """ + ) -> tuple[ + list[tuple[str, list[Union[None, str, list[float]]]]], # queries, params + list[tuple[int, str]], # idx, query_text pairs to embed + ]: + queries = [] + embedding_requests = [] + + for idx, (_, op) in enumerate(search_ops): + base_query = """ SELECT prefix, key, value, created_at, updated_at FROM store WHERE prefix LIKE %s """ params: list = [f"{_namespace_to_text(op.namespace_prefix)}%"] + needs_vector_search = False + + if op.query and self.embedding_config: + needs_vector_search = True + embedding_requests.append((idx, op.query)) + base_query = """ + SELECT s.prefix, s.key, s.value, s.created_at, s.updated_at, + 1 - (sv.embedding <=> %s::vector) as score + FROM store s + JOIN store_vectors sv ON s.prefix = sv.prefix AND s.key = sv.key + WHERE s.prefix LIKE %s + """ + params = [None, f"{_namespace_to_text(op.namespace_prefix)}%"] if op.filter: filter_conditions = [] for key, value in op.filter.items(): - if isinstance(value, list): - filter_conditions.append("value->%s @> %s::jsonb") - params.extend([key, json.dumps(value)]) + if isinstance(value, dict): + for op_name, val in value.items(): + condition, filter_params = self._get_filter_condition( + key, op_name, val + ) + filter_conditions.append(condition) + params.extend(filter_params) else: filter_conditions.append("value->%s = %s::jsonb") params.extend([key, json.dumps(value)]) - query += " AND " + " AND ".join(filter_conditions) - # Note: we will need to not do this if sim/keyword search - # is used - query += " ORDER BY updated_at DESC LIMIT %s OFFSET %s" + if filter_conditions: + base_query += " AND " + " AND ".join(filter_conditions) + + order_by = ( + "ORDER BY score DESC" + if needs_vector_search + else "ORDER BY updated_at DESC" + ) + base_query += f" {order_by} LIMIT %s OFFSET %s" params.extend([op.limit, op.offset]) + queries.append((base_query, params)) - queries.append((query, params)) - return queries + return queries, embedding_requests def _get_batch_list_namespaces_queries( self, @@ -222,10 +408,27 @@ class BasePostgresStore(Generic[C]): query += " ORDER BY truncated_prefix LIMIT %s OFFSET %s" params.extend([op.limit, op.offset]) - queries.append((query, params)) + queries.append((query, tuple(params))) return queries + def _get_filter_condition(self, key: str, op: str, value: Any) -> tuple[str, list]: + """Helper to generate filter conditions.""" + if op == "$eq": + return "value->%s = %s::jsonb", [key, json.dumps(value)] + elif op == "$gt": + return "value->>%s > %s", [key, str(value)] + elif op == "$gte": + return "value->>%s >= %s", [key, str(value)] + elif op == "$lt": + return "value->>%s < %s", [key, str(value)] + elif op == "$lte": + return "value->>%s <= %s", [key, str(value)] + elif op == "$ne": + return "value->%s != %s::jsonb", [key, json.dumps(value)] + else: + raise ValueError(f"Unsupported operator: {op}") + class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): __slots__ = ("_deserializer", "pipe", "lock", "supports_pipeline") @@ -238,6 +441,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): deserializer: Optional[ Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]] ] = None, + embedding: Optional[EmbeddingConfig] = None, ) -> None: super().__init__() self._deserializer = deserializer @@ -245,17 +449,24 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): self.pipe = pipe self.supports_pipeline = Capabilities().has_pipeline() self.lock = threading.Lock() + self.embedding_config = embedding + # TODO: Coerce embedding regular functions @classmethod @contextmanager def from_conn_string( - cls, conn_string: str, *, pipeline: bool = False + cls, + conn_string: str, + *, + pipeline: bool = False, + embedding: Optional[EmbeddingConfig] = None, ) -> Iterator["PostgresStore"]: """Create a new PostgresStore instance from a connection string. Args: conn_string (str): The Postgres connection info string. pipeline (bool): whether to use Pipeline + embedding (Optional[EmbeddingConfig]): The embedding config. Returns: PostgresStore: A new PostgresStore instance. @@ -265,9 +476,9 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): ) as conn: if pipeline: with conn.pipeline() as pipe: - yield cls(conn, pipe=pipe) + yield cls(conn, pipe=pipe, embedding=embedding) else: - yield cls(conn) + yield cls(conn, embedding=embedding) @contextmanager def _cursor(self, *, pipeline: bool = False) -> Iterator[Cursor[DictRow]]: @@ -363,7 +574,28 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): put_ops: Sequence[tuple[int, PutOp]], cur: Cursor[DictRow], ) -> None: - queries = self._get_batch_PUT_queries(put_ops) + queries, embedding_request = self._prepare_batch_PUT_queries(put_ops) + if embedding_request: + if self.embedding_config is None: + # Should not get here since the embedding config is required + # to return an embedding_request above + raise ValueError( + "Embedding configuration is required for vector operations " + f"(for semantic search). " + f"Please provide an EmbeddingConfig when initializing the {self.__class__.__name__}." + ) + query, txt_params = embedding_request + # Update the params to replace the raw text with the vectors + vectors = self.embedding_config["embed"].embed_documents( + [param[-1] for param in txt_params] + ) + queries.extend( + [ + (query, (ns, key, value, vector)) + for (ns, key, value, _), vector in zip(txt_params, vectors) + ] + ) + for query, params in queries: cur.execute(query, params) @@ -373,17 +605,28 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): results: list[Result], cur: Cursor[DictRow], ) -> None: - for (query, params), (idx, _) in zip( - self._get_batch_search_queries(search_ops), search_ops - ): + queries, embedding_requests = self._prepare_batch_search_queries(search_ops) + + if embedding_requests and self.embedding_config: + embeddings = self.embedding_config["embed"].embed_documents( + [query for _, query in embedding_requests] + ) + for (idx, _), embedding in zip(embedding_requests, embeddings): + queries[idx][1][0] = embedding + + for (idx, _), (query, params) in zip(search_ops, queries): cur.execute(query, params) rows = cast(list[Row], cur.fetchall()) - results[idx] = [ + items = [ _row_to_item( - _decode_ns_bytes(row["prefix"]), row, loader=self._deserializer + _decode_ns_bytes(row["prefix"]), + row, + loader=self._deserializer, + cls=SearchItem, ) for row in rows ] + results[idx] = items def _batch_list_namespaces_ops( self, @@ -424,11 +667,41 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]): ) """ ) + for v, migration in enumerate( self.MIGRATIONS[version + 1 :], start=version + 1 ): - cur.execute(migration) - cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,)) + if isinstance(migration, str): + sql = migration + else: + if migration.condition and not migration.condition(self): + continue + + sql = migration.sql + if migration.params: + params = { + k: v(self) if v is not None and callable(v) else v + for k, v in migration.params.items() + } + try: + sql = sql % params + except Exception as e: + logger.warning(f"Failed to format migration {v}: {e}") + if migration.condition == check_vector_available: + self.embedding_config = None + continue + + try: + cur.execute(sql) + cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,)) + except Exception as e: + logger.warning(f"Failed to run migration {v}: {e}") + if ( + not isinstance(migration, str) + and migration.condition == check_vector_available + ): + self.embedding_config = None + continue class Row(TypedDict): @@ -453,17 +726,32 @@ def _row_to_item( row: Row, *, loader: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]] = None, -) -> Item: - """Convert a row from the database into an Item.""" - loader = loader or _json_loads + cls: Union[Type[SearchItem], Type[Item]] = Item, +) -> Union[Item, SearchItem]: + """Convert a row from the database into an Item. + + Args: + namespace: Item namespace + row: Database row + loader: Optional value loader for non-dict values + cls: Item class to instantiate (Item or SearchItem) + """ val = row["value"] - return Item( - value=val if isinstance(val, dict) else loader(val), - key=row["key"], - namespace=namespace, - created_at=row["created_at"], - updated_at=row["updated_at"], - ) + if not isinstance(val, dict): + val = (loader or _json_loads)(val) + + kwargs = { + "key": row["key"], + "namespace": namespace, + "value": val, + "created_at": row["created_at"], + "updated_at": row["updated_at"], + } + + if cls is SearchItem and "score" in row: + kwargs["response_metadata"] = {"score": float(row["score"])} + + return cls(**kwargs) def _group_ops(ops: Iterable[Op]) -> tuple[dict[type, list[tuple[int, Op]]], int]: @@ -493,3 +781,157 @@ def _decode_ns_bytes(namespace: Union[str, bytes, list]) -> tuple[str, ...]: if isinstance(namespace, bytes): namespace = namespace.decode()[1:] return tuple(namespace.split(".")) + + +def _tokenize_path(path: str) -> list[str]: + """Tokenize a path into components. + + Handles: + - Simple paths: "field1.field2" + - Array indexing: "[0]", "[*]", "[-1]" + - Wildcards: "*" + - Multi-field selection: "{field1,field2}" + """ + if not path: + return [] + + tokens = [] + current: List[str] = [] + i = 0 + while i < len(path): + char = path[i] + + if char == "[": # Handle array index + if current: + tokens.append("".join(current)) + current = [] + bracket_count = 1 + index_chars = ["["] + i += 1 + while i < len(path) and bracket_count > 0: + if path[i] == "[": + bracket_count += 1 + elif path[i] == "]": + bracket_count -= 1 + index_chars.append(path[i]) + i += 1 + tokens.append("".join(index_chars)) + continue + + elif char == "{": # Handle multi-field selection + if current: + tokens.append("".join(current)) + current = [] + brace_count = 1 + field_chars = ["{"] + i += 1 + while i < len(path) and brace_count > 0: + if path[i] == "{": + brace_count += 1 + elif path[i] == "}": + brace_count -= 1 + field_chars.append(path[i]) + i += 1 + tokens.append("".join(field_chars)) + continue + + elif char == ".": + if current: + tokens.append("".join(current)) + current = [] + else: + current.append(char) + i += 1 + + if current: + tokens.append("".join(current)) + + return tokens + + +def _extract_text_by_path(obj: Any, path: str) -> list[str]: + """Extract text from an object using a path expression. + + Supports: + - Simple paths: "field1.field2" + - Array indexing: "[0]", "[*]", "[-1]" + - Wildcards: "*" + - Multi-field selection: "{field1,field2}" + - Nested paths in multi-field: "{field1,nested.field2}" + """ + if not path or path == "__root__": + return [json.dumps(obj, sort_keys=True)] + + def _extract_from_obj(obj: Any, tokens: list[str], pos: int) -> list[str]: + if pos >= len(tokens): + if isinstance(obj, (str, int, float, bool)): + return [str(obj)] + elif obj is None: + return [] + elif isinstance(obj, (list, dict)): + return [json.dumps(obj, sort_keys=True)] + return [] + + token = tokens[pos] + results = [] + + if token.startswith("[") and token.endswith("]"): + if not isinstance(obj, list): + return [] + + index = token[1:-1] + if index == "*": + for item in obj: + results.extend(_extract_from_obj(item, tokens, pos + 1)) + else: + try: + idx = int(index) + if idx < 0: + idx = len(obj) + idx + if 0 <= idx < len(obj): + results.extend(_extract_from_obj(obj[idx], tokens, pos + 1)) + except (ValueError, IndexError): + return [] + + elif token.startswith("{") and token.endswith("}"): + if not isinstance(obj, dict): + return [] + + fields = [f.strip() for f in token[1:-1].split(",")] + for field in fields: + nested_tokens = _tokenize_path(field) + if nested_tokens: + current_obj: Optional[dict] = obj + for nested_token in nested_tokens: + if ( + isinstance(current_obj, dict) + and nested_token in current_obj + ): + current_obj = current_obj[nested_token] + else: + current_obj = None + break + if current_obj is not None: + if isinstance(current_obj, (str, int, float, bool)): + results.append(str(current_obj)) + elif isinstance(current_obj, (list, dict)): + results.append(json.dumps(current_obj, sort_keys=True)) + + # Handle wildcard + elif token == "*": + if isinstance(obj, dict): + for value in obj.values(): + results.extend(_extract_from_obj(value, tokens, pos + 1)) + elif isinstance(obj, list): + for item in obj: + results.extend(_extract_from_obj(item, tokens, pos + 1)) + + # Handle regular field + else: + if isinstance(obj, dict) and token in obj: + results.extend(_extract_from_obj(obj[token], tokens, pos + 1)) + + return results + + tokens = _tokenize_path(path) + return _extract_from_obj(obj, tokens, 0) diff --git a/libs/checkpoint-postgres/tests/compose-postgres.yml b/libs/checkpoint-postgres/tests/compose-postgres.yml index a8a6c1e74..721784433 100644 --- a/libs/checkpoint-postgres/tests/compose-postgres.yml +++ b/libs/checkpoint-postgres/tests/compose-postgres.yml @@ -1,12 +1,13 @@ services: postgres-test: - image: postgres:${POSTGRES_VERSION:-16} + image: pgvector/pgvector:pg${POSTGRES_VERSION:-16} ports: - "5441:5432" environment: POSTGRES_DB: postgres POSTGRES_USER: postgres POSTGRES_PASSWORD: postgres + command: ["postgres", "-c", "shared_preload_libraries=vector"] healthcheck: test: pg_isready -U postgres start_period: 10s diff --git a/libs/checkpoint-postgres/tests/conftest.py b/libs/checkpoint-postgres/tests/conftest.py index 56d199812..bee32406b 100644 --- a/libs/checkpoint-postgres/tests/conftest.py +++ b/libs/checkpoint-postgres/tests/conftest.py @@ -4,6 +4,7 @@ import pytest from psycopg import AsyncConnection from psycopg.errors import UndefinedTable from psycopg.rows import DictRow, dict_row +from utils import CharacterEmbeddings # type: ignore DEFAULT_URI = "postgres://postgres:postgres@localhost:5441/postgres?sslmode=disable" @@ -31,3 +32,8 @@ async def clear_test_db(conn: AsyncConnection[DictRow]) -> None: await conn.execute("DELETE FROM store") except UndefinedTable: pass + + +@pytest.fixture +def fake_embeddings() -> CharacterEmbeddings: + return CharacterEmbeddings() diff --git a/libs/checkpoint-postgres/tests/test_async_store.py b/libs/checkpoint-postgres/tests/test_async_store.py index 8f9098eb2..e9ceb96a4 100644 --- a/libs/checkpoint-postgres/tests/test_async_store.py +++ b/libs/checkpoint-postgres/tests/test_async_store.py @@ -6,6 +6,7 @@ from typing import AsyncIterator import pytest from conftest import DEFAULT_URI # type: ignore from psycopg import AsyncConnection +from test_store import CharacterEmbeddings from langgraph.store.base import GetOp, Item, ListNamespacesOp, PutOp, SearchOp from langgraph.store.postgres import AsyncPostgresStore @@ -182,272 +183,179 @@ async def test_batch_list_namespaces_ops(store: AsyncPostgresStore) -> None: assert ("test", "namespace2") in results[0] -class TestAsyncPostgresStore: - @pytest.fixture(autouse=True) - async def setup(self) -> None: - async with AsyncPostgresStore.from_conn_string(DEFAULT_URI) as store: +@pytest.fixture +async def vector_store( + fake_embeddings: CharacterEmbeddings, +) -> AsyncIterator[AsyncPostgresStore]: + """Create a store with vector search enabled.""" + if sys.version_info < (3, 10): + pytest.skip("Async Postgres tests require Python 3.10+") + + database = f"test_{uuid.uuid4().hex[:16]}" + uri_parts = DEFAULT_URI.split("/") + uri_base = "/".join(uri_parts[:-1]) + query_params = "" + if "?" in uri_parts[-1]: + db_name, query_params = uri_parts[-1].split("?", 1) + query_params = "?" + query_params + + conn_string = f"{uri_base}/{database}{query_params}" + admin_conn_string = DEFAULT_URI + + async with await AsyncConnection.connect( + admin_conn_string, autocommit=True + ) as conn: + await conn.execute(f"CREATE DATABASE {database}") + try: + async with AsyncPostgresStore.from_conn_string( + conn_string, + embedding={"dims": fake_embeddings.dims, "embed": fake_embeddings}, + ) as store: await store.setup() + yield store + finally: + async with await AsyncConnection.connect( + admin_conn_string, autocommit=True + ) as conn: + await conn.execute(f"DROP DATABASE {database}") - async def test_basic_store_ops(self) -> None: - async with AsyncPostgresStore.from_conn_string(DEFAULT_URI) as store: - namespace = ("test", "documents") - item_id = "doc1" - item_value = {"title": "Test Document", "content": "Hello, World!"} - await store.aput(namespace, item_id, item_value) - item = await store.aget(namespace, item_id) +async def test_vector_store_initialization( + vector_store: AsyncPostgresStore, fake_embeddings: CharacterEmbeddings +) -> None: + """Test store initialization with embedding config.""" + assert vector_store.embedding_config is not None + assert vector_store.embedding_config["dims"] == fake_embeddings.dims + assert vector_store.embedding_config["embed"] == fake_embeddings - assert item - assert item.namespace == namespace - assert item.key == item_id - assert item.value == item_value - updated_value = { - "title": "Updated Test Document", - "content": "Hello, LangGraph!", - } - await store.aput(namespace, item_id, updated_value) - updated_item = await store.aget(namespace, item_id) +async def test_vector_insert_with_auto_embedding( + vector_store: AsyncPostgresStore, +) -> None: + """Test inserting items that get auto-embedded.""" + docs = [ + ("doc1", {"text": "short text"}), + ("doc2", {"text": "longer text document"}), + ("doc3", {"text": "longest text document here"}), + ("doc4", {"description": "text in description field"}), + ("doc5", {"content": "text in content field"}), + ("doc6", {"body": "text in body field"}), + ] - assert updated_item.value == updated_value - assert updated_item.updated_at > item.updated_at - different_namespace = ("test", "other_documents") - item_in_different_namespace = await store.aget(different_namespace, item_id) - assert item_in_different_namespace is None + for key, value in docs: + await vector_store.aput(("test",), key, value) - new_item_id = "doc2" - new_item_value = {"title": "Another Document", "content": "Greetings!"} - await store.aput(namespace, new_item_id, new_item_value) + results = await vector_store.asearch(("test",), query="long text") + assert len(results) > 0 - search_results = await store.asearch(["test"], limit=10) - items = search_results - assert len(items) == 2 - assert any(item.key == item_id for item in items) - assert any(item.key == new_item_id for item in items) + doc_order = [r.key for r in results] + assert "doc2" in doc_order + assert "doc3" in doc_order - namespaces = await store.alist_namespaces(prefix=["test"]) - assert ("test", "documents") in namespaces - await store.adelete(namespace, item_id) - await store.adelete(namespace, new_item_id) - deleted_item = await store.aget(namespace, item_id) - assert deleted_item is None +async def test_vector_update_with_embedding(vector_store: AsyncPostgresStore) -> None: + """Test that updating items properly updates their embeddings.""" + await vector_store.aput(("test",), "doc1", {"text": "initial text about cats"}) + await vector_store.aput(("test",), "doc2", {"text": "something about dogs"}) + await vector_store.aput(("test",), "doc3", {"text": "text about birds"}) - deleted_item = await store.aget(namespace, new_item_id) - assert deleted_item is None + results_initial = await vector_store.asearch(("test",), query="cats") + assert len(results_initial) > 0 + assert results_initial[0].key == "doc1" + initial_score = results_initial[0].response_metadata["score"] - empty_search_results = await store.asearch(["test"], limit=10) - assert len(empty_search_results) == 0 + await vector_store.aput(("test",), "doc1", {"text": "new text about dogs"}) - async def test_list_namespaces(self) -> None: - async with AsyncPostgresStore.from_conn_string(DEFAULT_URI) as store: - test_pref = str(uuid.uuid4()) - test_namespaces = [ - (test_pref, "test", "documents", "public", test_pref), - (test_pref, "test", "documents", "private", test_pref), - (test_pref, "test", "images", "public", test_pref), - (test_pref, "test", "images", "private", test_pref), - (test_pref, "prod", "documents", "public", test_pref), - ( - test_pref, - "prod", - "documents", - "some", - "nesting", - "public", - test_pref, - ), - (test_pref, "prod", "documents", "private", test_pref), - ] + results_after = await vector_store.asearch(("test",), query="cats") + after_score = next( + (r.response_metadata["score"] for r in results_after if r.key == "doc1"), 0.0 + ) + assert after_score < initial_score - for namespace in test_namespaces: - await store.aput(namespace, "dummy", {"content": "dummy"}) + results_new = await vector_store.asearch(("test",), query="dogs") + for r in results_new: + if r.key == "doc1": + assert r.response_metadata["score"] > after_score - prefix_result = await store.alist_namespaces(prefix=[test_pref, "test"]) - assert len(prefix_result) == 4 - assert all([ns[1] == "test" for ns in prefix_result]) - specific_prefix_result = await store.alist_namespaces( - prefix=[test_pref, "test", "documents"] - ) - assert len(specific_prefix_result) == 2 - assert all( - [ns[1:3] == ("test", "documents") for ns in specific_prefix_result] - ) +async def test_vector_search_with_filters(vector_store: AsyncPostgresStore) -> None: + """Test combining vector search with filters.""" + docs = [ + ("doc1", {"text": "red apple", "color": "red", "score": 4.5}), + ("doc2", {"text": "red car", "color": "red", "score": 3.0}), + ("doc3", {"text": "green apple", "color": "green", "score": 4.0}), + ("doc4", {"text": "blue car", "color": "blue", "score": 3.5}), + ] - suffix_result = await store.alist_namespaces(suffix=["public", test_pref]) - assert len(suffix_result) == 4 - assert all(ns[-2] == "public" for ns in suffix_result) + for key, value in docs: + await vector_store.aput(("test",), key, value) - prefix_suffix_result = await store.alist_namespaces( - prefix=[test_pref, "test"], suffix=["public", test_pref] - ) - assert len(prefix_suffix_result) == 2 - assert all( - ns[1] == "test" and ns[-2] == "public" for ns in prefix_suffix_result - ) + results = await vector_store.asearch( + ("test",), query="apple", filter={"color": "red"} + ) + assert len(results) == 2 + assert results[0].key == "doc1" - wildcard_prefix_result = await store.alist_namespaces( - prefix=[test_pref, "*", "documents"] - ) - assert len(wildcard_prefix_result) == 5 - assert all(ns[2] == "documents" for ns in wildcard_prefix_result) + results = await vector_store.asearch( + ("test",), query="car", filter={"color": "red"} + ) + assert len(results) == 2 + assert results[0].key == "doc2" - wildcard_suffix_result = await store.alist_namespaces( - suffix=["*", "public", test_pref] - ) - assert len(wildcard_suffix_result) == 4 - assert all(ns[-2] == "public" for ns in wildcard_suffix_result) - wildcard_single = await store.alist_namespaces( - suffix=["some", "*", "public", test_pref] - ) - assert len(wildcard_single) == 1 - assert wildcard_single[0] == ( - test_pref, - "prod", - "documents", - "some", - "nesting", - "public", - test_pref, - ) + results = await vector_store.asearch( + ("test",), query="bbbbluuu", filter={"score": {"$gt": 3.2}} + ) + assert len(results) == 3 + assert results[0].key == "doc4" - max_depth_result = await store.alist_namespaces(max_depth=3) - assert all([len(ns) <= 3 for ns in max_depth_result]) - max_depth_result = await store.alist_namespaces( - max_depth=4, prefix=[test_pref, "*", "documents"] - ) - assert ( - len(set(tuple(res) for res in max_depth_result)) - == len(max_depth_result) - == 5 - ) + results = await vector_store.asearch( + ("test",), query="apple", filter={"score": {"$gte": 4.0}, "color": "green"} + ) + assert len(results) == 1 + assert results[0].key == "doc3" - limit_result = await store.alist_namespaces(prefix=[test_pref], limit=3) - assert len(limit_result) == 3 - offset_result = await store.alist_namespaces(prefix=[test_pref], offset=3) - assert len(offset_result) == len(test_namespaces) - 3 +async def test_vector_search_pagination(vector_store: AsyncPostgresStore) -> None: + """Test pagination with vector search.""" + for i in range(5): + await vector_store.aput( + ("test",), f"doc{i}", {"text": f"test document number {i}"} + ) - empty_prefix_result = await store.alist_namespaces(prefix=[test_pref]) - assert len(empty_prefix_result) == len(test_namespaces) - assert set(tuple(ns) for ns in empty_prefix_result) == set( - tuple(ns) for ns in test_namespaces - ) + results_page1 = await vector_store.asearch(("test",), query="test", limit=2) + results_page2 = await vector_store.asearch( + ("test",), query="test", limit=2, offset=2 + ) - for namespace in test_namespaces: - await store.adelete(namespace, "dummy") + assert len(results_page1) == 2 + assert len(results_page2) == 2 + assert results_page1[0].key != results_page2[0].key - async def test_search(self): - async with AsyncPostgresStore.from_conn_string(DEFAULT_URI) as store: - test_namespaces = [ - ("test_search", "documents", "user1"), - ("test_search", "documents", "user2"), - ("test_search", "reports", "department1"), - ("test_search", "reports", "department2"), - ] - test_items = [ - {"title": "Doc 1", "author": "John Doe", "tags": ["important"]}, - {"title": "Doc 2", "author": "Jane Smith", "tags": ["draft"]}, - {"title": "Report A", "author": "John Doe", "tags": ["final"]}, - {"title": "Report B", "author": "Alice Johnson", "tags": ["draft"]}, - ] - empty = await store.asearch( - ( - "scoped", - "assistant_id", - "shared", - "6c5356f6-63ab-4158-868d-cd9fd14c736e", - ), - limit=10, - offset=0, - ) - assert len(empty) == 0 + all_results = await vector_store.asearch(("test",), query="test", limit=10) + assert len(all_results) == 5 - for namespace, item in zip(test_namespaces, test_items): - await store.aput(namespace, f"item_{namespace[-1]}", item) - docs_result = await store.asearch(["test_search", "documents"]) - assert len(docs_result) == 2 - assert all([item.namespace[1] == "documents" for item in docs_result]), [ - item.namespace for item in docs_result - ] +async def test_vector_search_edge_cases(vector_store: AsyncPostgresStore) -> None: + """Test edge cases in vector search.""" + await vector_store.aput(("test",), "doc1", {"text": "test document"}) - reports_result = await store.asearch(["test_search", "reports"]) - assert len(reports_result) == 2 - assert all(item.namespace[1] == "reports" for item in reports_result) + perfect_match = await vector_store.asearch(("test",), query="text test document") + perfect_score = perfect_match[0].response_metadata["score"] - limited_result = await store.asearch(["test_search"], limit=2) - assert len(limited_result) == 2 - offset_result = await store.asearch(["test_search"]) - assert len(offset_result) == 4 + results = await vector_store.asearch(("test",), query="") + assert len(results) == 1 + assert "score" not in results[0].response_metadata - offset_result = await store.asearch(["test_search"], offset=2) - assert len(offset_result) == 2 - assert all(item not in limited_result for item in offset_result) + results = await vector_store.asearch(("test",), query=None) + assert len(results) == 1 + assert "score" not in results[0].response_metadata - john_doe_result = await store.asearch( - ["test_search"], filter={"author": "John Doe"} - ) - assert len(john_doe_result) == 2 - assert all(item.value["author"] == "John Doe" for item in john_doe_result) + long_query = "foo " * 100 + results = await vector_store.asearch(("test",), query=long_query) + assert len(results) == 1 + assert results[0].response_metadata["score"] < perfect_score - draft_result = await store.asearch( - ["test_search"], filter={"tags": ["draft"]} - ) - assert len(draft_result) == 2 - assert all("draft" in item.value["tags"] for item in draft_result) - - page1 = await store.asearch(["test_search"], limit=2, offset=0) - page2 = await store.asearch(["test_search"], limit=2, offset=2) - all_items = page1 + page2 - assert len(all_items) == 4 - assert len(set(item.key for item in all_items)) == 4 - empty = await store.asearch( - ( - "scoped", - "assistant_id", - "shared", - "again", - "maybe", - "some-long", - "6be5cb0e-2eb4-42e6-bb6b-fba3c269db25", - ), - limit=10, - offset=0, - ) - assert len(empty) == 0 - - # Test with a namespace beginning with a number (like a UUID) - uuid_namespace = (str(uuid.uuid4()), "documents") - uuid_item_id = "uuid_doc" - uuid_item_value = { - "title": "UUID Document", - "content": "This document has a UUID namespace.", - } - - # Insert the item with the UUID namespace - await store.aput(uuid_namespace, uuid_item_id, uuid_item_value) - - # Retrieve the item to verify it was stored correctly - retrieved_item = await store.aget(uuid_namespace, uuid_item_id) - assert retrieved_item is not None - assert retrieved_item.namespace == uuid_namespace - assert retrieved_item.key == uuid_item_id - assert retrieved_item.value == uuid_item_value - - # Search for the item using the UUID namespace - search_result = await store.asearch([uuid_namespace[0]]) - assert len(search_result) == 1 - assert search_result[0].key == uuid_item_id - assert search_result[0].value == uuid_item_value - - # Clean up: delete the item with the UUID namespace - await store.adelete(uuid_namespace, uuid_item_id) - - # Verify the item was deleted - deleted_item = await store.aget(uuid_namespace, uuid_item_id) - assert deleted_item is None - - for namespace in test_namespaces: - await store.adelete(namespace, f"item_{namespace[-1]}") + special_query = "test!@#$%^&*()" + results = await vector_store.asearch(("test",), query=special_query) + assert len(results) == 1 + assert results[0].response_metadata["score"] < perfect_score diff --git a/libs/checkpoint-postgres/tests/test_store.py b/libs/checkpoint-postgres/tests/test_store.py index 402eaedc5..69d401f27 100644 --- a/libs/checkpoint-postgres/tests/test_store.py +++ b/libs/checkpoint-postgres/tests/test_store.py @@ -1,11 +1,14 @@ # type: ignore +import json from uuid import uuid4 import pytest from conftest import DEFAULT_URI # type: ignore +from langchain_core.embeddings import Embeddings from psycopg import Connection from psycopg_pool import ConnectionPool +from utils import CharacterEmbeddings from langgraph.store.base import ( GetOp, @@ -16,6 +19,7 @@ from langgraph.store.base import ( SearchOp, ) from langgraph.store.postgres import PostgresStore +from langgraph.store.postgres.base import _extract_text_by_path @pytest.fixture(scope="function", params=["default", "pipe", "pool"]) @@ -342,3 +346,229 @@ class TestPostgresStore: # Cleanup for namespace, key, _ in test_data: store.delete(namespace, key) + + +@pytest.fixture +def vector_store(fake_embeddings: Embeddings) -> PostgresStore: + """Create a store with vector search enabled.""" + database = f"test_{uuid4().hex[:16]}" + uri_parts = DEFAULT_URI.split("/") + uri_base = "/".join(uri_parts[:-1]) + query_params = "" + if "?" in uri_parts[-1]: + db_name, query_params = uri_parts[-1].split("?", 1) + query_params = "?" + query_params + + conn_string = f"{uri_base}/{database}{query_params}" + admin_conn_string = DEFAULT_URI + + with Connection.connect(admin_conn_string, autocommit=True) as conn: + conn.execute(f"CREATE DATABASE {database}") + try: + with PostgresStore.from_conn_string( + conn_string, + embedding={"dims": fake_embeddings.dims, "embed": fake_embeddings}, + ) as store: + store.setup() + yield store + finally: + with Connection.connect(admin_conn_string, autocommit=True) as conn: + conn.execute(f"DROP DATABASE {database}") + + +def test_vector_store_initialization( + vector_store: PostgresStore, fake_embeddings: CharacterEmbeddings +) -> None: + """Test store initialization with embedding config.""" + # Store should be initialized with embedding config + assert vector_store.embedding_config is not None + assert vector_store.embedding_config["dims"] == fake_embeddings.dims + assert vector_store.embedding_config["embed"] == fake_embeddings + + +def test_vector_insert_with_auto_embedding(vector_store: PostgresStore) -> None: + """Test inserting items that get auto-embedded.""" + docs = [ + ("doc1", {"text": "short text"}), + ("doc2", {"text": "longer text document"}), + ("doc3", {"text": "longest text document here"}), + ("doc4", {"description": "text in description field"}), + ("doc5", {"content": "text in content field"}), + ("doc6", {"body": "text in body field"}), + ] + + for key, value in docs: + vector_store.put(("test",), key, value) + + results = vector_store.search(("test",), query="long text") + assert len(results) > 0 + + doc_order = [r.key for r in results] + assert "doc2" in doc_order + assert "doc3" in doc_order + + +def test_vector_update_with_embedding(vector_store: PostgresStore) -> None: + """Test that updating items properly updates their embeddings.""" + vector_store.put(("test",), "doc1", {"text": "initial text about cats"}) + vector_store.put(("test",), "doc2", {"text": "something about dogs"}) + vector_store.put(("test",), "doc3", {"text": "text about birds"}) + + results_initial = vector_store.search(("test",), query="cats") + assert len(results_initial) > 0 + assert results_initial[0].key == "doc1" + initial_score = results_initial[0].response_metadata["score"] + + vector_store.put(("test",), "doc1", {"text": "new text about dogs"}) + + results_after = vector_store.search(("test",), query="cats") + after_score = next( + (r.response_metadata["score"] for r in results_after if r.key == "doc1"), 0.0 + ) + assert after_score < initial_score + + results_new = vector_store.search(("test",), query="dogs") + for r in results_new: + if r.key == "doc1": + assert r.response_metadata["score"] > after_score + + +def test_vector_search_with_filters(vector_store: PostgresStore) -> None: + """Test combining vector search with filters.""" + # Insert test documents + docs = [ + ("doc1", {"text": "red apple", "color": "red", "score": 4.5}), + ("doc2", {"text": "red car", "color": "red", "score": 3.0}), + ("doc3", {"text": "green apple", "color": "green", "score": 4.0}), + ("doc4", {"text": "blue car", "color": "blue", "score": 3.5}), + ] + + for key, value in docs: + vector_store.put(("test",), key, value) + + results = vector_store.search(("test",), query="apple", filter={"color": "red"}) + assert len(results) == 2 + assert results[0].key == "doc1" + + results = vector_store.search(("test",), query="car", filter={"color": "red"}) + assert len(results) == 2 + assert results[0].key == "doc2" + + results = vector_store.search( + ("test",), query="bbbbluuu", filter={"score": {"$gt": 3.2}} + ) + assert len(results) == 3 + assert results[0].key == "doc4" + + # Multiple filters + results = vector_store.search( + ("test",), query="apple", filter={"score": {"$gte": 4.0}, "color": "green"} + ) + assert len(results) == 1 + assert results[0].key == "doc3" + + +def test_vector_search_pagination(vector_store: PostgresStore) -> None: + """Test pagination with vector search.""" + # Insert multiple similar documents + for i in range(5): + vector_store.put(("test",), f"doc{i}", {"text": f"test document number {i}"}) + + # Test with different page sizes + results_page1 = vector_store.search(("test",), query="test", limit=2) + results_page2 = vector_store.search(("test",), query="test", limit=2, offset=2) + + assert len(results_page1) == 2 + assert len(results_page2) == 2 + assert results_page1[0].key != results_page2[0].key + + # Get all results + all_results = vector_store.search(("test",), query="test", limit=10) + assert len(all_results) == 5 + + +def test_vector_search_edge_cases(vector_store: PostgresStore) -> None: + """Test edge cases in vector search.""" + vector_store.put(("test",), "doc1", {"text": "test document"}) + + results = vector_store.search(("test",), query="") + assert len(results) == 1 + + results = vector_store.search(("test",), query=None) + assert len(results) == 1 + + long_query = "test " * 100 + results = vector_store.search(("test",), query=long_query) + assert len(results) == 1 + + special_query = "test!@#$%^&*()" + results = vector_store.search(("test",), query=special_query) + assert len(results) == 1 + + +def test_extract_text_by_path(): + nested_data = { + "name": "test", + "info": { + "age": 25, + "tags": ["a", "b", "c"], + "metadata": {"created": "2024-01-01", "updated": "2024-01-02"}, + }, + "items": [ + {"id": 1, "value": "first", "tags": ["x", "y"]}, + {"id": 2, "value": "second", "tags": ["y", "z"]}, + {"id": 3, "value": "third", "tags": ["z", "w"]}, + ], + "empty": None, + "zeros": [0, 0.0, "0"], + "empty_list": [], + "empty_dict": {}, + } + + assert _extract_text_by_path(nested_data, "__root__") == [ + json.dumps(nested_data, sort_keys=True) + ] + + assert _extract_text_by_path(nested_data, "name") == ["test"] + assert _extract_text_by_path(nested_data, "info.age") == ["25"] + + assert _extract_text_by_path(nested_data, "info.metadata.created") == ["2024-01-01"] + + assert _extract_text_by_path(nested_data, "items[0].value") == ["first"] + assert _extract_text_by_path(nested_data, "items[-1].value") == ["third"] + assert _extract_text_by_path(nested_data, "items[1].tags[0]") == ["y"] + + values = _extract_text_by_path(nested_data, "items[*].value") + assert set(values) == {"first", "second", "third"} + + metadata_dates = _extract_text_by_path(nested_data, "info.metadata.*") + assert set(metadata_dates) == {"2024-01-01", "2024-01-02"} + name_and_age = _extract_text_by_path(nested_data, "{name,info.age}") + assert set(name_and_age) == {"test", "25"} + + item_fields = _extract_text_by_path(nested_data, "items[*].{id,value}") + assert set(item_fields) == {"1", "2", "3", "first", "second", "third"} + + all_tags = _extract_text_by_path(nested_data, "items[*].tags[*]") + assert set(all_tags) == {"x", "y", "z", "w"} + + assert _extract_text_by_path(None, "any.path") == [] + assert _extract_text_by_path({}, "any.path") == [] + assert _extract_text_by_path(nested_data, "") == [ + json.dumps(nested_data, sort_keys=True) + ] + assert _extract_text_by_path(nested_data, "nonexistent") == [] + assert _extract_text_by_path(nested_data, "items[99].value") == [] + assert _extract_text_by_path(nested_data, "items[*].nonexistent") == [] + + assert _extract_text_by_path(nested_data, "empty") == [] + assert _extract_text_by_path(nested_data, "empty_list") == ["[]"] + assert _extract_text_by_path(nested_data, "empty_dict") == ["{}"] + + zeros = _extract_text_by_path(nested_data, "zeros[*]") + assert set(zeros) == {"0", "0.0", "0"} + + assert _extract_text_by_path(nested_data, "items[].value") == [] + assert _extract_text_by_path(nested_data, "items[abc].value") == [] + assert _extract_text_by_path(nested_data, "{unclosed") == [] + assert _extract_text_by_path(nested_data, "nested[{invalid}]") == [] diff --git a/libs/checkpoint-postgres/tests/utils.py b/libs/checkpoint-postgres/tests/utils.py new file mode 100644 index 000000000..3e045e8ae --- /dev/null +++ b/libs/checkpoint-postgres/tests/utils.py @@ -0,0 +1,63 @@ +import math +import random +from collections import Counter +from typing import Any, Optional + +from langchain_core.embeddings import Embeddings + + +class CharacterEmbeddings(Embeddings): + """Simple character-frequency based embeddings using random projections.""" + + def __init__(self, dims: int = 50, seed: int = 42): + """Initialize with embedding dimensions and random seed.""" + self._rng = random.Random(seed) + self._char_to_idx: dict[str, int] = {} + self._projection: Optional[list[list[float]]] = None + self.dims = dims + + def _ensure_projection_matrix(self, texts: list[str]) -> None: + """Lazily initialize character mapping and projection matrix.""" + if self._projection is None: + chars = sorted(set("".join(texts))) + self._char_to_idx = {c: i for i, c in enumerate(chars)} + self._projection = [ + [self._rng.gauss(0, 1 / math.sqrt(self.dims)) for _ in range(self.dims)] + for _ in range(len(chars)) + ] + + def _embed_one(self, text: str) -> list[float]: + """Embed a single text.""" + counts = Counter(text) + char_vec = [0.0] * len(self._char_to_idx) + + for char, count in counts.items(): + if char in self._char_to_idx: + char_vec[self._char_to_idx[char]] = count + + total = sum(char_vec) + if total > 0: + char_vec = [v / total for v in char_vec] + embedding = [ + sum(a * b for a, b in zip(char_vec, proj)) + for proj in zip(*self._projection) + ] + + norm = math.sqrt(sum(x * x for x in embedding)) + if norm > 0: + embedding = [x / norm for x in embedding] + + return embedding + + def embed_documents(self, texts: list[str]) -> list[list[float]]: + """Embed a list of documents.""" + self._ensure_projection_matrix(texts) + return [self._embed_one(text) for text in texts] + + def embed_query(self, text: str) -> list[float]: + """Embed a query string.""" + self._ensure_projection_matrix([text]) + return self._embed_one(text) + + def __eq__(self, other: Any) -> bool: + return isinstance(other, CharacterEmbeddings) and self.dims == other.dims diff --git a/libs/checkpoint/langgraph/store/base/__init__.py b/libs/checkpoint/langgraph/store/base/__init__.py index 098462339..7f5073f78 100644 --- a/libs/checkpoint/langgraph/store/base/__init__.py +++ b/libs/checkpoint/langgraph/store/base/__init__.py @@ -6,7 +6,7 @@ scoped to user IDs, assistant IDs, or other arbitrary namespaces. from abc import ABC, abstractmethod from datetime import datetime -from typing import Any, Iterable, Literal, NamedTuple, Optional, Union, cast +from typing import Any, Iterable, Literal, NamedTuple, Optional, TypedDict, Union, cast class Item: @@ -73,6 +73,52 @@ class Item: } +class ResponseMetadata(TypedDict, total=False): + """Additional metadata about the response/result.""" + + score: float + """Relevance/similarity score if from a ranked operation.""" + + +class SearchItem(Item): + """Represents a result item with additional response metadata.""" + + __slots__ = "response_metadata" + + def __init__( + self, + namespace: tuple[str, ...], + key: str, + value: dict[str, Any], + created_at: datetime, + updated_at: datetime, + response_metadata: Optional[ResponseMetadata] = None, + ) -> None: + """Initialize a result item. + + Args: + namespace: Hierarchical path to the item. + key: Unique identifier within the namespace. + value: The stored value. + created_at: When the item was first created. + updated_at: When the item was last updated. + response_metadata: Optional metadata about the response/result. + """ + super().__init__( + value=value, + key=key, + namespace=namespace, + created_at=created_at, + updated_at=updated_at, + ) + self.response_metadata = response_metadata or {} + + def dict(self) -> dict: + result = super().dict() + result["response_metadata"] = self.response_metadata + return result + + class GetOp(NamedTuple): """Operation to retrieve an item by namespace and key.""" @@ -93,6 +139,8 @@ class SearchOp(NamedTuple): """Maximum number of items to return.""" offset: int = 0 """Number of items to skip before returning results.""" + query: Optional[str] = None + """The search query for natural language search.""" class PutOp(NamedTuple): @@ -231,6 +279,7 @@ class BaseStore(ABC): namespace_prefix: tuple[str, ...], /, *, + query: Optional[str] = None, filter: Optional[dict[str, Any]] = None, limit: int = 10, offset: int = 0, @@ -239,6 +288,7 @@ class BaseStore(ABC): Args: namespace_prefix: Hierarchical path prefix to search within. + query: Optional query for natural language search. filter: Key-value pairs to filter results. limit: Maximum number of items to return. offset: Number of items to skip before returning results. @@ -246,7 +296,7 @@ class BaseStore(ABC): Returns: List of items matching the search criteria. """ - return self.batch([SearchOp(namespace_prefix, filter, limit, offset)])[0] + return self.batch([SearchOp(namespace_prefix, filter, limit, offset, query)])[0] def put(self, namespace: tuple[str, ...], key: str, value: dict[str, Any]) -> None: """Store or update an item. @@ -336,6 +386,7 @@ class BaseStore(ABC): namespace_prefix: tuple[str, ...], /, *, + query: Optional[str] = None, filter: Optional[dict[str, Any]] = None, limit: int = 10, offset: int = 0, @@ -344,6 +395,7 @@ class BaseStore(ABC): Args: namespace_prefix: Hierarchical path prefix to search within. + query: Optional query for natural language search. filter: Key-value pairs to filter results. limit: Maximum number of items to return. offset: Number of items to skip before returning results. @@ -351,9 +403,11 @@ class BaseStore(ABC): Returns: List of items matching the search criteria. """ - return (await self.abatch([SearchOp(namespace_prefix, filter, limit, offset)]))[ - 0 - ] + return ( + await self.abatch( + [SearchOp(namespace_prefix, filter, limit, offset, query)] + ) + )[0] async def aput( self, namespace: tuple[str, ...], key: str, value: dict[str, Any] diff --git a/libs/checkpoint/langgraph/store/base/batch.py b/libs/checkpoint/langgraph/store/base/batch.py index 079888222..545ab1d1c 100644 --- a/libs/checkpoint/langgraph/store/base/batch.py +++ b/libs/checkpoint/langgraph/store/base/batch.py @@ -40,12 +40,13 @@ class AsyncBatchedBaseStore(BaseStore): namespace_prefix: tuple[str, ...], /, *, + query: Optional[str] = None, filter: Optional[dict[str, Any]] = None, limit: int = 10, offset: int = 0, ) -> list[Item]: fut = self._loop.create_future() - self._aqueue[fut] = SearchOp(namespace_prefix, filter, limit, offset) + self._aqueue[fut] = SearchOp(namespace_prefix, filter, limit, offset, query) return await fut async def aput(