mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 13:35:09 +02:00
Move embedding config to base
This commit is contained in:
@@ -2,7 +2,7 @@ import asyncio
|
||||
import logging
|
||||
from collections.abc import AsyncIterator, Iterable, Sequence
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, Callable, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Callable, Optional, Union, cast
|
||||
|
||||
import orjson
|
||||
from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities
|
||||
@@ -19,18 +19,22 @@ from langgraph.store.base import (
|
||||
Result,
|
||||
SearchItem,
|
||||
SearchOp,
|
||||
ensure_embeddings,
|
||||
)
|
||||
from langgraph.store.base.batch import AsyncBatchedBaseStore
|
||||
from langgraph.store.postgres.base import (
|
||||
BasePostgresStore,
|
||||
EmbeddingConfig,
|
||||
PoolConfig,
|
||||
PostgresEmbeddingConfig,
|
||||
Row,
|
||||
_decode_ns_bytes,
|
||||
_group_ops,
|
||||
_row_to_item,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.embeddings import Embeddings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -51,7 +55,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
deserializer: Optional[
|
||||
Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]
|
||||
] = None,
|
||||
embedding: Optional[EmbeddingConfig] = None,
|
||||
embedding: Optional[PostgresEmbeddingConfig] = None,
|
||||
) -> None:
|
||||
if isinstance(conn, AsyncConnectionPool) and pipe is not None:
|
||||
raise ValueError(
|
||||
@@ -65,6 +69,13 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
self.loop = asyncio.get_running_loop()
|
||||
self.supports_pipeline = Capabilities().has_pipeline()
|
||||
self.embedding_config = embedding
|
||||
if self.embedding_config:
|
||||
self.embeddings: Optional[Embeddings] = ensure_embeddings(
|
||||
self.embedding_config.get("embed"),
|
||||
aembed=self.embedding_config.get("aembed"),
|
||||
)
|
||||
else:
|
||||
self.embeddings = None
|
||||
|
||||
async def abatch(self, ops: Iterable[Op]) -> list[Result]:
|
||||
grouped_ops, num_ops = _group_ops(ops)
|
||||
@@ -142,7 +153,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
) -> None:
|
||||
queries, embedding_request = self._prepare_batch_PUT_queries(put_ops)
|
||||
if embedding_request:
|
||||
if self.embedding_config is None:
|
||||
if self.embeddings is None:
|
||||
# Should not get here since the embedding config is required
|
||||
# to return an embedding_request above
|
||||
raise ValueError(
|
||||
@@ -152,7 +163,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
)
|
||||
query, txt_params = embedding_request
|
||||
# Update the params to replace the raw text with the vectors
|
||||
vectors = await self.embedding_config["embed"].aembed_documents(
|
||||
vectors = await self.embeddings.aembed_documents(
|
||||
[param[-1] for param in txt_params]
|
||||
)
|
||||
queries.extend(
|
||||
@@ -173,8 +184,8 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
) -> None:
|
||||
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(
|
||||
if embedding_requests and self.embeddings:
|
||||
embeddings = await self.embeddings.aembed_documents(
|
||||
[query for _, query in embedding_requests]
|
||||
)
|
||||
for (idx, _), embedding in zip(embedding_requests, embeddings):
|
||||
@@ -261,7 +272,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
*,
|
||||
pipeline: bool = False,
|
||||
pool_config: Optional[PoolConfig] = None,
|
||||
embedding: Optional[EmbeddingConfig] = None,
|
||||
embedding: Optional[PostgresEmbeddingConfig] = None,
|
||||
) -> AsyncIterator["AsyncPostgresStore"]:
|
||||
"""Create a new AsyncPostgresStore instance from a connection string.
|
||||
|
||||
@@ -271,7 +282,7 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
pool_config (Optional[PoolConfig]): Configuration for the connection pool.
|
||||
If provided, will create a connection pool and use it instead of a single connection.
|
||||
This overrides the `pipeline` argument.
|
||||
embedding (Optional[EmbeddingConfig]): The embedding config.
|
||||
embedding (Optional[PostgresEmbeddingConfig]): The embedding config.
|
||||
|
||||
Returns:
|
||||
AsyncPostgresStore: A new AsyncPostgresStore instance.
|
||||
|
||||
@@ -6,10 +6,20 @@ from collections import defaultdict
|
||||
from collections.abc import Iterable, Iterator, Sequence
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime
|
||||
from typing import Any, Callable, Generic, NamedTuple, Optional, TypeVar, Union, cast
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
Generic,
|
||||
Literal,
|
||||
NamedTuple,
|
||||
Optional,
|
||||
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
|
||||
@@ -21,6 +31,7 @@ from langgraph.checkpoint.postgres import _ainternal as _ainternal
|
||||
from langgraph.checkpoint.postgres import _internal as _pg_internal
|
||||
from langgraph.store.base import (
|
||||
BaseStore,
|
||||
EmbeddingConfig,
|
||||
GetOp,
|
||||
Item,
|
||||
ListNamespacesOp,
|
||||
@@ -29,37 +40,15 @@ from langgraph.store.base import (
|
||||
Result,
|
||||
SearchItem,
|
||||
SearchOp,
|
||||
ensure_embeddings,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.embeddings import Embeddings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
|
||||
class Migration(NamedTuple):
|
||||
"""A database migration with optional conditions and parameters."""
|
||||
|
||||
@@ -73,6 +62,48 @@ def _embedding_requested(store: Any) -> bool:
|
||||
return bool(store.embedding_config)
|
||||
|
||||
|
||||
def _get_vector_type_ops(store: Any) -> str:
|
||||
"""Get the vector type operator class based on config."""
|
||||
if not store.embedding_config:
|
||||
return "vector_cosine_ops"
|
||||
|
||||
config = cast(PostgresEmbeddingConfig, store.embedding_config)
|
||||
index_config = config.get(
|
||||
"index_config", BasePostgresStore._get_default_index_config()
|
||||
)
|
||||
vector_type = index_config.get("vector_type", "vector")
|
||||
distance_type = config.get("distance_type", "cosine")
|
||||
|
||||
# For regular vectors
|
||||
type_prefix = {"vector": "vector", "halfvec": "halfvec"}[vector_type]
|
||||
|
||||
if distance_type not in ("l2", "inner_product", "cosine"):
|
||||
raise ValueError(
|
||||
f"Vector type {vector_type} only supports 'l2', 'inner_product', or 'cosine' distance, got {distance_type}"
|
||||
)
|
||||
|
||||
distance_suffix = {
|
||||
"l2": "l2_ops",
|
||||
"inner_product": "ip_ops",
|
||||
"cosine": "cosine_ops",
|
||||
}[distance_type]
|
||||
|
||||
return f"{type_prefix}_{distance_suffix}"
|
||||
|
||||
|
||||
def _get_index_params(store: Any) -> tuple[str, dict[str, Any]]:
|
||||
"""Get the index type and configuration based on config."""
|
||||
if not store.embedding_config:
|
||||
return "hnsw", {}
|
||||
|
||||
config = cast(PostgresEmbeddingConfig, store.embedding_config)
|
||||
default_config = BasePostgresStore._get_default_index_config()
|
||||
index_config = config.get("index_config", default_config).copy()
|
||||
kind = index_config.pop("kind", "hnsw")
|
||||
index_config.pop("vector_type", None)
|
||||
return kind, index_config
|
||||
|
||||
|
||||
MIGRATIONS: Sequence[Union[str, Migration]] = [
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS store (
|
||||
@@ -101,7 +132,7 @@ CREATE TABLE IF NOT EXISTS store_vectors (
|
||||
prefix text NOT NULL,
|
||||
key text NOT NULL,
|
||||
field_name text NOT NULL,
|
||||
embedding vector(%(dims)s),
|
||||
embedding %(vector_type)s(%(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),
|
||||
@@ -109,17 +140,36 @@ CREATE TABLE IF NOT EXISTS store_vectors (
|
||||
);
|
||||
""",
|
||||
condition=_embedding_requested,
|
||||
params={"dims": lambda store: store.embedding_config["dims"]},
|
||||
params={
|
||||
"dims": lambda store: store.embedding_config["dims"],
|
||||
"vector_type": lambda store: (
|
||||
cast(PostgresEmbeddingConfig, store.embedding_config)
|
||||
.get("index_config", {})
|
||||
.get("vector_type", "vector")
|
||||
),
|
||||
},
|
||||
),
|
||||
Migration(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS store_vectors_embedding_idx ON store_vectors
|
||||
USING ivfflat (embedding vector_cosine_ops);
|
||||
USING %(index_type)s (embedding %(ops)s)%(index_params)s;
|
||||
""",
|
||||
condition=_embedding_requested,
|
||||
params={
|
||||
"index_type": lambda store: _get_index_params(store)[0],
|
||||
"ops": lambda store: _get_vector_type_ops(store),
|
||||
"index_params": lambda store: (
|
||||
" WITH ("
|
||||
+ ", ".join(f"{k}={v}" for k, v in _get_index_params(store)[1].items())
|
||||
+ ")"
|
||||
if _get_index_params(store)[1]
|
||||
else ""
|
||||
),
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
C = TypeVar("C", bound=Union[_pg_internal.Conn, _ainternal.Conn])
|
||||
|
||||
|
||||
@@ -148,11 +198,76 @@ class PoolConfig(TypedDict, total=False):
|
||||
"""
|
||||
|
||||
|
||||
class IndexConfig(TypedDict, total=False):
|
||||
"""Configuration for vector index in PostgreSQL store."""
|
||||
|
||||
kind: Literal["hnsw", "ivfflat"]
|
||||
"""Type of index to use: 'hnsw' for Hierarchical Navigable Small World, or 'ivfflat' for Inverted File Flat."""
|
||||
vector_type: Literal["vector", "halfvec"]
|
||||
"""Type of vector storage to use.
|
||||
Options:
|
||||
- 'vector': Regular vectors (default)
|
||||
- 'halfvec': Half-precision vectors for reduced memory usage
|
||||
"""
|
||||
|
||||
|
||||
class HNSWConfig(IndexConfig, total=False):
|
||||
"""Configuration for HNSW (Hierarchical Navigable Small World) index."""
|
||||
|
||||
kind: Literal["hnsw"] # type: ignore[misc]
|
||||
m: int
|
||||
"""Maximum number of connections per layer. Default is 16."""
|
||||
ef_construction: int
|
||||
"""Size of dynamic candidate list for index construction. Default is 64."""
|
||||
|
||||
|
||||
class IVFFlatConfig(IndexConfig, total=False):
|
||||
"""IVFFlat index divides vectors into lists, and then searches a subset of those lists that are closest to the query vector. It has faster build times and uses less memory than HNSW, but has lower query performance (in terms of speed-recall tradeoff).
|
||||
|
||||
Three keys to achieving good recall are:
|
||||
1. Create the index after the table has some data
|
||||
2. Choose an appropriate number of lists - a good place to start is rows / 1000 for up to 1M rows and sqrt(rows) for over 1M rows
|
||||
3. When querying, specify an appropriate number of probes (higher is better for recall, lower is better for speed) - a good place to start is sqrt(lists)
|
||||
"""
|
||||
|
||||
kind: Literal["ivfflat"] # type: ignore[misc]
|
||||
nlist: int
|
||||
"""Number of inverted lists (clusters) for IVF index.
|
||||
|
||||
Determines the number of clusters used in the index structure.
|
||||
Higher values can improve search speed but increase index size and build time.
|
||||
Typically set to the square root of the number of vectors in the index.
|
||||
"""
|
||||
|
||||
|
||||
class PostgresEmbeddingConfig(EmbeddingConfig, total=False):
|
||||
"""Configuration for vector embeddings in PostgreSQL store with pgvector-specific options.
|
||||
|
||||
Extends EmbeddingConfig with additional configuration for pgvector index and vector types.
|
||||
"""
|
||||
|
||||
index_config: Union[HNSWConfig, IVFFlatConfig]
|
||||
"""Specific configuration for the chosen index type (HNSW or IVF Flat)."""
|
||||
distance_type: Literal["l2", "inner_product", "cosine"]
|
||||
"""Distance metric to use for vector similarity search:
|
||||
- 'l2': Euclidean distance
|
||||
- 'inner_product': Dot product
|
||||
- 'cosine': Cosine similarity
|
||||
"""
|
||||
|
||||
|
||||
class BasePostgresStore(Generic[C]):
|
||||
MIGRATIONS = MIGRATIONS
|
||||
conn: C
|
||||
_deserializer: Optional[Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]]
|
||||
embedding_config: Optional[EmbeddingConfig]
|
||||
embedding_config: Optional[PostgresEmbeddingConfig]
|
||||
|
||||
@staticmethod
|
||||
def _get_default_index_config() -> IndexConfig:
|
||||
return HNSWConfig(
|
||||
kind="hnsw",
|
||||
vector_type="vector",
|
||||
)
|
||||
|
||||
def _get_batch_GET_ops_queries(
|
||||
self,
|
||||
@@ -234,7 +349,6 @@ class BasePostgresStore(Generic[C]):
|
||||
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)
|
||||
@@ -292,9 +406,26 @@ class BasePostgresStore(Generic[C]):
|
||||
if op.query and self.embedding_config:
|
||||
needs_vector_search = True
|
||||
embedding_requests.append((idx, op.query))
|
||||
base_query = """
|
||||
|
||||
_, score_expr = _get_distance_operator(self)
|
||||
vector_type = (
|
||||
cast(PostgresEmbeddingConfig, self.embedding_config)
|
||||
.get("index_config", self._get_default_index_config())
|
||||
.get("vector_type", "vector")
|
||||
)
|
||||
|
||||
# For hamming distance, we need the vector dimension for normalization
|
||||
if (
|
||||
vector_type == "bit"
|
||||
and self.embedding_config.get("distance_type") == "hamming"
|
||||
):
|
||||
score_expr = score_expr % ("%s", self.embedding_config["dims"])
|
||||
else:
|
||||
score_expr = score_expr % ("%s", vector_type)
|
||||
|
||||
base_query = f"""
|
||||
SELECT s.prefix, s.key, s.value, s.created_at, s.updated_at,
|
||||
1 - (sv.embedding <=> %s::vector) as score
|
||||
{score_expr} as score
|
||||
FROM store s
|
||||
JOIN store_vectors sv ON s.prefix = sv.prefix AND s.key = sv.key
|
||||
WHERE s.prefix LIKE %s
|
||||
@@ -412,7 +543,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
deserializer: Optional[
|
||||
Callable[[Union[bytes, orjson.Fragment]], dict[str, Any]]
|
||||
] = None,
|
||||
embedding: Optional[EmbeddingConfig] = None,
|
||||
embedding: Optional[PostgresEmbeddingConfig] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self._deserializer = deserializer
|
||||
@@ -421,6 +552,13 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
self.supports_pipeline = Capabilities().has_pipeline()
|
||||
self.lock = threading.Lock()
|
||||
self.embedding_config = embedding
|
||||
if self.embedding_config:
|
||||
self.embeddings: Optional[Embeddings] = ensure_embeddings(
|
||||
self.embedding_config.get("embed"),
|
||||
aembed=self.embedding_config.get("aembed"),
|
||||
)
|
||||
else:
|
||||
self.embeddings = None
|
||||
# TODO: Coerce embedding regular functions
|
||||
|
||||
@classmethod
|
||||
@@ -431,7 +569,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
*,
|
||||
pipeline: bool = False,
|
||||
pool_config: Optional[PoolConfig] = None,
|
||||
embedding: Optional[EmbeddingConfig] = None,
|
||||
embedding: Optional[PostgresEmbeddingConfig] = None,
|
||||
) -> Iterator["PostgresStore"]:
|
||||
"""Create a new PostgresStore instance from a connection string.
|
||||
|
||||
@@ -441,7 +579,7 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
pool_config (Optional[PoolArgs]): Configuration for the connection pool.
|
||||
If provided, will create a connection pool and use it instead of a single connection.
|
||||
This overrides the `pipeline` argument.
|
||||
embedding (Optional[EmbeddingConfig]): The embedding config.
|
||||
embedding (Optional[PostgresEmbeddingConfig]): The embedding config.
|
||||
|
||||
Returns:
|
||||
PostgresStore: A new PostgresStore instance.
|
||||
@@ -574,19 +712,20 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
) -> None:
|
||||
queries, embedding_request = self._prepare_batch_PUT_queries(put_ops)
|
||||
if embedding_request:
|
||||
if self.embedding_config is None:
|
||||
if self.embeddings 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__}."
|
||||
f"Please provide an Embeddings 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(
|
||||
vectors = self.embeddings.embed_documents(
|
||||
[param[-1] for param in txt_params]
|
||||
)
|
||||
|
||||
queries.extend(
|
||||
[
|
||||
(query, (ns, key, value, vector))
|
||||
@@ -605,8 +744,8 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
|
||||
) -> None:
|
||||
queries, embedding_requests = self._prepare_batch_search_queries(search_ops)
|
||||
|
||||
if embedding_requests and self.embedding_config:
|
||||
embeddings = self.embedding_config["embed"].embed_documents(
|
||||
if embedding_requests and self.embeddings:
|
||||
embeddings = self.embeddings.embed_documents(
|
||||
[query for _, query in embedding_requests]
|
||||
)
|
||||
for (idx, _), embedding in zip(embedding_requests, embeddings):
|
||||
@@ -818,7 +957,7 @@ def _tokenize_path(path: str) -> list[str]:
|
||||
tokens.append("".join(field_chars))
|
||||
continue
|
||||
|
||||
elif char == ".":
|
||||
elif char == ".": # Handle regular field
|
||||
if current:
|
||||
tokens.append("".join(current))
|
||||
current = []
|
||||
@@ -918,3 +1057,19 @@ def _extract_text_by_path(obj: Any, path: str) -> list[str]:
|
||||
|
||||
tokens = _tokenize_path(path)
|
||||
return _extract_from_obj(obj, tokens, 0)
|
||||
|
||||
|
||||
def _get_distance_operator(store: Any) -> tuple[str, str]:
|
||||
"""Get the distance operator and score expression based on config."""
|
||||
if not store.embedding_config:
|
||||
return "<=>", "1 - (sv.embedding <=> %s::vector)"
|
||||
|
||||
config = cast(PostgresEmbeddingConfig, store.embedding_config)
|
||||
distance_type = config.get("distance_type", "cosine")
|
||||
|
||||
if distance_type == "l2":
|
||||
return "<->", "1 - (sv.embedding <-> %s::%s)"
|
||||
elif distance_type == "inner_product":
|
||||
return "<#>", "-(sv.embedding <#> %s::%s)"
|
||||
else: # cosine
|
||||
return "<=>", "1 - (sv.embedding <=> %s::%s)"
|
||||
|
||||
@@ -36,4 +36,8 @@ async def clear_test_db(conn: AsyncConnection[DictRow]) -> None:
|
||||
|
||||
@pytest.fixture
|
||||
def fake_embeddings() -> CharacterEmbeddings:
|
||||
return CharacterEmbeddings()
|
||||
return CharacterEmbeddings(dims=500)
|
||||
|
||||
|
||||
INDEX_TYPES = ["hnsw", "ivfflat"]
|
||||
VECTOR_TYPES = ["vector", "halfvec"]
|
||||
|
||||
@@ -4,9 +4,14 @@ import uuid
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
import pytest
|
||||
from conftest import DEFAULT_URI # type: ignore
|
||||
from conftest import (
|
||||
DEFAULT_URI, # type: ignore
|
||||
INDEX_TYPES,
|
||||
VECTOR_TYPES,
|
||||
CharacterEmbeddings,
|
||||
)
|
||||
from langchain_core.embeddings import Embeddings
|
||||
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,8 +187,22 @@ async def test_batch_list_namespaces_ops(store: AsyncPostgresStore) -> None:
|
||||
assert ("test", "namespace2") in results[0]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@pytest.fixture(
|
||||
scope="function",
|
||||
params=[
|
||||
(index_type, vector_type, distance_type)
|
||||
for index_type in INDEX_TYPES
|
||||
for vector_type in VECTOR_TYPES
|
||||
for distance_type in (
|
||||
(["hamming"] if index_type == "ivfflat" else ["hamming", "jaccard"])
|
||||
if vector_type == "bit"
|
||||
else ["l2", "inner_product", "cosine"]
|
||||
)
|
||||
],
|
||||
ids=lambda p: f"{p[0]}_{p[1]}_{p[2]}",
|
||||
)
|
||||
async def vector_store(
|
||||
request,
|
||||
fake_embeddings: CharacterEmbeddings,
|
||||
) -> AsyncIterator[AsyncPostgresStore]:
|
||||
"""Create a store with vector search enabled."""
|
||||
@@ -201,6 +220,17 @@ async def vector_store(
|
||||
conn_string = f"{uri_base}/{database}{query_params}"
|
||||
admin_conn_string = DEFAULT_URI
|
||||
|
||||
index_type, vector_type, distance_type = request.param
|
||||
embedding_config = {
|
||||
"dims": fake_embeddings.dims,
|
||||
"embed": fake_embeddings,
|
||||
"index_config": {
|
||||
"kind": index_type,
|
||||
"vector_type": vector_type,
|
||||
},
|
||||
"distance_type": distance_type,
|
||||
}
|
||||
|
||||
async with await AsyncConnection.connect(
|
||||
admin_conn_string, autocommit=True
|
||||
) as conn:
|
||||
@@ -208,7 +238,7 @@ async def vector_store(
|
||||
try:
|
||||
async with AsyncPostgresStore.from_conn_string(
|
||||
conn_string,
|
||||
embedding={"dims": fake_embeddings.dims, "embed": fake_embeddings},
|
||||
embedding=embedding_config,
|
||||
) as store:
|
||||
await store.setup()
|
||||
yield store
|
||||
@@ -225,7 +255,8 @@ async def test_vector_store_initialization(
|
||||
"""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
|
||||
if isinstance(vector_store.embedding_config["embed"], Embeddings):
|
||||
assert vector_store.embedding_config["embed"] == fake_embeddings
|
||||
|
||||
|
||||
async def test_vector_insert_with_auto_embedding(
|
||||
@@ -254,24 +285,24 @@ async def test_vector_insert_with_auto_embedding(
|
||||
|
||||
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",), "doc1", {"text": "zany zebra Xerxes"})
|
||||
await vector_store.aput(("test",), "doc2", {"text": "something about dogs"})
|
||||
await vector_store.aput(("test",), "doc3", {"text": "text about birds"})
|
||||
|
||||
results_initial = await vector_store.asearch(("test",), query="cats")
|
||||
results_initial = await vector_store.asearch(("test",), query="Zany Xerxes")
|
||||
assert len(results_initial) > 0
|
||||
assert results_initial[0].key == "doc1"
|
||||
initial_score = results_initial[0].response_metadata["score"]
|
||||
|
||||
await vector_store.aput(("test",), "doc1", {"text": "new text about dogs"})
|
||||
|
||||
results_after = await vector_store.asearch(("test",), query="cats")
|
||||
results_after = await vector_store.asearch(("test",), query="Zany Xerxes")
|
||||
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 = await vector_store.asearch(("test",), query="dogs")
|
||||
results_new = await vector_store.asearch(("test",), query="new text about dogs")
|
||||
for r in results_new:
|
||||
if r.key == "doc1":
|
||||
assert r.response_metadata["score"] > after_score
|
||||
|
||||
@@ -4,10 +4,14 @@ import json
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from conftest import DEFAULT_URI # type: ignore
|
||||
from conftest import (
|
||||
DEFAULT_URI, # type: ignore
|
||||
INDEX_TYPES,
|
||||
VECTOR_TYPES,
|
||||
CharacterEmbeddings,
|
||||
)
|
||||
from langchain_core.embeddings import Embeddings
|
||||
from psycopg import Connection
|
||||
from utils import CharacterEmbeddings
|
||||
|
||||
from langgraph.store.base import (
|
||||
GetOp,
|
||||
@@ -346,8 +350,21 @@ class TestPostgresStore:
|
||||
store.delete(namespace, key)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def vector_store(fake_embeddings: Embeddings) -> PostgresStore:
|
||||
@pytest.fixture(
|
||||
scope="function",
|
||||
params=[
|
||||
(index_type, vector_type, distance_type)
|
||||
for index_type in INDEX_TYPES
|
||||
for vector_type in VECTOR_TYPES
|
||||
for distance_type in (
|
||||
(["hamming"] if index_type == "ivfflat" else ["hamming", "jaccard"])
|
||||
if vector_type == "bit"
|
||||
else ["l2", "inner_product", "cosine"]
|
||||
)
|
||||
],
|
||||
ids=lambda p: f"{p[0]}_{p[1]}_{p[2]}",
|
||||
)
|
||||
def vector_store(request, fake_embeddings: Embeddings) -> PostgresStore:
|
||||
"""Create a store with vector search enabled."""
|
||||
database = f"test_{uuid4().hex[:16]}"
|
||||
uri_parts = DEFAULT_URI.split("/")
|
||||
@@ -360,12 +377,23 @@ def vector_store(fake_embeddings: Embeddings) -> PostgresStore:
|
||||
conn_string = f"{uri_base}/{database}{query_params}"
|
||||
admin_conn_string = DEFAULT_URI
|
||||
|
||||
index_type, vector_type, distance_type = request.param
|
||||
embedding_config = {
|
||||
"dims": fake_embeddings.dims,
|
||||
"embed": fake_embeddings,
|
||||
"index_config": {
|
||||
"kind": index_type,
|
||||
"vector_type": vector_type,
|
||||
},
|
||||
"distance_type": distance_type,
|
||||
}
|
||||
|
||||
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},
|
||||
embedding=embedding_config,
|
||||
) as store:
|
||||
store.setup()
|
||||
yield store
|
||||
@@ -408,24 +436,24 @@ def test_vector_insert_with_auto_embedding(vector_store: PostgresStore) -> None:
|
||||
|
||||
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",), "doc1", {"text": "zany zebra Xerxes"})
|
||||
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")
|
||||
results_initial = vector_store.search(("test",), query="Zany Xerxes")
|
||||
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")
|
||||
results_after = vector_store.search(("test",), query="Zany Xerxes")
|
||||
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")
|
||||
results_new = vector_store.search(("test",), query="new text about dogs")
|
||||
for r in results_new:
|
||||
if r.key == "doc1":
|
||||
assert r.response_metadata["score"] > after_score
|
||||
|
||||
@@ -6,7 +6,11 @@ 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, TypedDict, Union, cast
|
||||
from typing import ( Any, Iterable, Literal, NamedTuple,
|
||||
Optional, TypedDict, Union, cast)
|
||||
|
||||
from langchain_core.embeddings import Embeddings
|
||||
from langgraph.store.base._embed import AEmbeddingsFunc, EmbeddingsFunc, ensure_embeddings
|
||||
|
||||
|
||||
class Item:
|
||||
@@ -229,6 +233,38 @@ def _validate_namespace(namespace: tuple[str, ...]) -> None:
|
||||
)
|
||||
|
||||
|
||||
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: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc]
|
||||
"""Optional function to generate embeddings from text."""
|
||||
aembed: Optional[AEmbeddingsFunc]
|
||||
"""Optional asynchronous function to generate embeddings from text.
|
||||
|
||||
Provide for asynchronous embedding generation if you do not provide
|
||||
an Embeddings object.
|
||||
"""
|
||||
|
||||
text_fields: Optional[list[str]]
|
||||
"""Fields to extract text from for embedding generation.
|
||||
|
||||
Defaults to ["__root__"], which embeds the json object as a whole.
|
||||
"""
|
||||
|
||||
|
||||
class BaseStore(ABC):
|
||||
"""Abstract base class for persistent key-value stores.
|
||||
|
||||
@@ -298,7 +334,13 @@ class BaseStore(ABC):
|
||||
"""
|
||||
return self.batch([SearchOp(namespace_prefix, filter, limit, offset, query)])[0]
|
||||
|
||||
def put(self, namespace: tuple[str, ...], key: str, value: dict[str, Any]) -> None:
|
||||
def put(
|
||||
self,
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
value: dict[str, Any],
|
||||
index: Optional[bool] = None,
|
||||
) -> None:
|
||||
"""Store or update an item.
|
||||
|
||||
Args:
|
||||
@@ -410,7 +452,11 @@ class BaseStore(ABC):
|
||||
)[0]
|
||||
|
||||
async def aput(
|
||||
self, namespace: tuple[str, ...], key: str, value: dict[str, Any]
|
||||
self,
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
value: dict[str, Any],
|
||||
index: Optional[bool] = None,
|
||||
) -> None:
|
||||
"""Asynchronously store or update an item.
|
||||
|
||||
@@ -481,3 +527,18 @@ class BaseStore(ABC):
|
||||
offset=offset,
|
||||
)
|
||||
return (await self.abatch([op]))[0]
|
||||
|
||||
__all__ = [
|
||||
"BaseStore",
|
||||
"Item",
|
||||
"Op",
|
||||
"PutOp",
|
||||
"GetOp",
|
||||
"SearchOp",
|
||||
"ListNamespacesOp",
|
||||
"MatchCondition",
|
||||
"NameSpacePath",
|
||||
"NamespaceMatchType",
|
||||
"Embeddings",
|
||||
"ensure_embeddings",
|
||||
]
|
||||
@@ -0,0 +1,205 @@
|
||||
"""Utilities for working with embedding functions and LangChain's Embeddings interface.
|
||||
|
||||
This module provides tools to wrap arbitrary embedding functions (both sync and async)
|
||||
into LangChain's Embeddings interface. This enables using custom embedding functions
|
||||
with LangChain-compatible tools while maintaining support for both synchronous and
|
||||
asynchronous operations.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import (Any, Awaitable, Callable, List, Optional, Sequence,
|
||||
TypeGuard, Union)
|
||||
|
||||
from langchain_core.embeddings import Embeddings
|
||||
|
||||
EmbeddingsFunc = Callable[[Sequence[str]], list[list[float]]]
|
||||
"""Type for synchronous embedding functions.
|
||||
|
||||
The function should take a sequence of strings and return a list of embeddings,
|
||||
where each embedding is a list of floats. The dimensionality of the embeddings
|
||||
should be consistent for all inputs.
|
||||
"""
|
||||
|
||||
AEmbeddingsFunc = Callable[[Sequence[str]], Awaitable[list[list[float]]]]
|
||||
"""Type for asynchronous embedding functions.
|
||||
|
||||
Similar to EmbeddingsFunc, but returns an awaitable that resolves to the embeddings.
|
||||
"""
|
||||
|
||||
|
||||
def ensure_embeddings(
|
||||
embed: Union[Embeddings, EmbeddingsFunc, AEmbeddingsFunc, None],
|
||||
*,
|
||||
aembed: Optional[AEmbeddingsFunc] = None
|
||||
) -> Embeddings:
|
||||
"""Ensure that an embedding function conforms to LangChain's Embeddings interface.
|
||||
|
||||
This function wraps arbitrary embedding functions to make them compatible with
|
||||
LangChain's Embeddings interface. It handles both synchronous and asynchronous
|
||||
functions.
|
||||
|
||||
Args:
|
||||
embed: Either an existing Embeddings instance, or a function that converts
|
||||
text to embeddings. If the function is async, it will be used for both
|
||||
sync and async operations.
|
||||
aembed: Optional async function for embeddings. If provided, it will be used
|
||||
for async operations while the sync function is used for sync operations.
|
||||
Must be None if embed is async.
|
||||
|
||||
Returns:
|
||||
An Embeddings instance that wraps the provided function(s).
|
||||
|
||||
Example:
|
||||
>>> def my_embed_fn(texts): return [[0.1, 0.2] for _ in texts]
|
||||
>>> async def my_async_fn(texts): return [[0.1, 0.2] for _ in texts]
|
||||
>>> # Wrap a sync function
|
||||
>>> embeddings = ensure_embeddings(my_embed_fn)
|
||||
>>> # Wrap an async function
|
||||
>>> embeddings = ensure_embeddings(my_async_fn)
|
||||
>>> # Provide both sync and async implementations
|
||||
>>> embeddings = ensure_embeddings(my_embed_fn, aembed=my_async_fn)
|
||||
"""
|
||||
if embed is None and aembed is None:
|
||||
raise ValueError("embed or aembed must be provided")
|
||||
if isinstance(embed, Embeddings):
|
||||
return embed
|
||||
return EmbeddingsLambda(embed, afunc=aembed)
|
||||
|
||||
|
||||
class EmbeddingsLambda(Embeddings):
|
||||
"""Wrapper to convert embedding functions into LangChain's Embeddings interface.
|
||||
|
||||
This class allows arbitrary embedding functions to be used with LangChain-compatible
|
||||
tools. It supports both synchronous and asynchronous operations, and can be
|
||||
initialized with either:
|
||||
1. A synchronous function for both sync/async operations
|
||||
2. An async function for both sync/async operations
|
||||
3. Both sync and async functions for their respective operations
|
||||
|
||||
The embedding functions should convert text into fixed-dimensional vectors that
|
||||
capture the semantic meaning of the text.
|
||||
|
||||
Args:
|
||||
func: Function that converts text to embeddings. Can be sync or async.
|
||||
If async, it will be used for both sync and async operations.
|
||||
afunc: Optional async function for embeddings. If provided, it will be used
|
||||
for async operations while func is used for sync operations.
|
||||
Must be None if func is async.
|
||||
|
||||
Example:
|
||||
>>> def my_embed_fn(texts):
|
||||
... # Return 2D embeddings for each text
|
||||
... return [[0.1, 0.2] for _ in texts]
|
||||
>>> embeddings = EmbeddingsLambda(my_embed_fn)
|
||||
>>> result = embeddings.embed_query("hello") # Returns [0.1, 0.2]
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
func: Union[EmbeddingsFunc, AEmbeddingsFunc, None],
|
||||
afunc: Optional[AEmbeddingsFunc] = None,
|
||||
) -> None:
|
||||
if _is_async_callable(func):
|
||||
if afunc is not None:
|
||||
raise ValueError(
|
||||
"afunc must be None if func is async. The async func will be used for both sync and async operations."
|
||||
)
|
||||
self.afunc = func
|
||||
else:
|
||||
self.func = func
|
||||
if afunc is not None:
|
||||
self.afunc = afunc
|
||||
|
||||
def embed_documents(self, texts: List[str]) -> List[List[float]]:
|
||||
"""Embed a list of texts into vectors.
|
||||
|
||||
Args:
|
||||
texts: List of texts to convert to embeddings.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per input text. Each embedding is a list of floats.
|
||||
|
||||
Raises:
|
||||
ValueError: If the instance was initialized with only an async function.
|
||||
"""
|
||||
if not hasattr(self, "func"):
|
||||
raise ValueError(
|
||||
"EmbeddingsLambda was initialized with an async function but no sync function. "
|
||||
"Use aembed_documents for async operation or provide a sync function."
|
||||
)
|
||||
return self.func(texts)
|
||||
|
||||
def embed_query(self, text: str) -> List[float]:
|
||||
"""Embed a single piece of text.
|
||||
|
||||
Args:
|
||||
text: Text to convert to an embedding.
|
||||
|
||||
Returns:
|
||||
Embedding vector as a list of floats.
|
||||
|
||||
Note:
|
||||
This is equivalent to calling embed_documents with a single text
|
||||
and taking the first result.
|
||||
"""
|
||||
return self.embed_documents([text])[0]
|
||||
|
||||
async def aembed_documents(self, texts: List[str]) -> List[List[float]]:
|
||||
"""Asynchronously embed a list of texts into vectors.
|
||||
|
||||
Args:
|
||||
texts: List of texts to convert to embeddings.
|
||||
|
||||
Returns:
|
||||
List of embeddings, one per input text. Each embedding is a list of floats.
|
||||
|
||||
Note:
|
||||
If no async function was provided, this falls back to the sync implementation.
|
||||
"""
|
||||
if not hasattr(self, "afunc"):
|
||||
return await super().aembed_documents(texts)
|
||||
return await self.afunc(texts)
|
||||
|
||||
async def aembed_query(self, text: str) -> List[float]:
|
||||
"""Asynchronously embed a single piece of text.
|
||||
|
||||
Args:
|
||||
text: Text to convert to an embedding.
|
||||
|
||||
Returns:
|
||||
Embedding vector as a list of floats.
|
||||
|
||||
Note:
|
||||
This is equivalent to calling aembed_documents with a single text
|
||||
and taking the first result.
|
||||
"""
|
||||
if not hasattr(self, "afunc"):
|
||||
return await super().aembed_query(text)
|
||||
return (await self.afunc([text]))[0]
|
||||
|
||||
|
||||
def _is_async_callable(
|
||||
func: Any,
|
||||
) -> TypeGuard[Callable[..., Awaitable]]:
|
||||
"""Check if a function is async.
|
||||
|
||||
This includes both async def functions and classes with async __call__ methods.
|
||||
|
||||
Args:
|
||||
func: Function or callable object to check.
|
||||
|
||||
Returns:
|
||||
True if the function is async, False otherwise.
|
||||
"""
|
||||
return (
|
||||
asyncio.iscoroutinefunction(func)
|
||||
or hasattr(func, "__call__") # noqa: B004
|
||||
and asyncio.iscoroutinefunction(func.__call__)
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ensure_embeddings",
|
||||
"EmbeddingsFunc",
|
||||
"AEmbeddingsFunc",
|
||||
]
|
||||
Reference in New Issue
Block a user