Move embedding config to base

This commit is contained in:
William Fu-Hinthorn
2024-11-25 16:30:29 -08:00
parent b1fe1718a1
commit 9c25bd1225
7 changed files with 570 additions and 75 deletions
@@ -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)"
+5 -1
View File
@@ -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
+37 -9
View File
@@ -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",
]