This commit is contained in:
William Fu-Hinthorn
2024-11-26 21:30:41 -08:00
10 changed files with 657 additions and 383 deletions
@@ -2,7 +2,7 @@ import asyncio
import logging
from collections.abc import AsyncIterator, Iterable, Sequence
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Any, Callable, Optional, Union, cast
from typing import Any, Callable, Optional, Union, cast
import orjson
from psycopg import AsyncConnection, AsyncCursor, AsyncPipeline, Capabilities
@@ -18,8 +18,6 @@ from langgraph.store.base import (
PutOp,
Result,
SearchOp,
ensure_embeddings,
tokenize_path,
)
from langgraph.store.base.batch import AsyncBatchedBaseStore
from langgraph.store.postgres.base import (
@@ -28,14 +26,12 @@ from langgraph.store.postgres.base import (
PostgresEmbeddingConfig,
Row,
_decode_ns_bytes,
_ensure_embedding_config,
_group_ops,
_row_to_item,
_row_to_search_item,
)
if TYPE_CHECKING:
from langchain_core.embeddings import Embeddings
logger = logging.getLogger(__name__)
@@ -71,15 +67,10 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
self.supports_pipeline = Capabilities().has_pipeline()
self.embedding_config = embedding
if self.embedding_config:
self.embedding_config = self.embedding_config.copy()
self.embedding_config["__tokenized_fields"] = [
(p, tokenize_path(p)) if p != "__root__" else (p, p)
for p in (self.embedding_config.get("text_fields") or ["__root__"])
]
self.embeddings: Optional[Embeddings] = ensure_embeddings(
self.embedding_config.get("embed"),
aembed=self.embedding_config.get("aembed"),
self.embeddings, self.embedding_config = _ensure_embedding_config(
self.embedding_config
)
else:
self.embeddings = None
@@ -96,179 +87,6 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
return results
async def _execute_batch(
self,
grouped_ops: dict,
results: list[Result],
conn: AsyncConnection[DictRow],
) -> None:
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]),
results,
cur,
)
if SearchOp in grouped_ops:
await self._batch_search_ops(
cast(Sequence[tuple[int, SearchOp]], grouped_ops[SearchOp]),
results,
cur,
)
if ListNamespacesOp in grouped_ops:
await self._batch_list_namespaces_ops(
cast(
Sequence[tuple[int, ListNamespacesOp]],
grouped_ops[ListNamespacesOp],
),
results,
cur,
)
if PutOp in grouped_ops:
await self._batch_put_ops(
cast(Sequence[tuple[int, PutOp]], grouped_ops[PutOp]),
cur,
)
async def _batch_get_ops(
self,
get_ops: Sequence[tuple[int, GetOp]],
results: list[Result],
cur: AsyncCursor[DictRow],
) -> None:
for query, params, namespace, items in self._get_batch_GET_ops_queries(get_ops):
await cur.execute(query, params)
rows = cast(list[Row], await cur.fetchall())
key_to_row = {row["key"]: row for row in rows}
for idx, key in items:
row = key_to_row.get(key)
if row:
results[idx] = _row_to_item(
namespace, row, loader=self._deserializer
)
else:
results[idx] = None
async def _batch_put_ops(
self,
put_ops: Sequence[tuple[int, PutOp]],
cur: AsyncCursor[DictRow],
) -> None:
queries, embedding_request = self._prepare_batch_PUT_queries(put_ops)
if embedding_request:
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__}."
)
query, txt_params = embedding_request
# Update the params to replace the raw text with the vectors
vectors = await self.embeddings.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)
async def _batch_search_ops(
self,
search_ops: Sequence[tuple[int, SearchOp]],
results: list[Result],
cur: AsyncCursor[DictRow],
) -> None:
queries, embedding_requests = self._prepare_batch_search_queries(search_ops)
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):
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_search_item(
_decode_ns_bytes(row["prefix"]),
row,
loader=self._deserializer,
)
for row in rows
]
results[idx] = items
async def _batch_list_namespaces_ops(
self,
list_ops: Sequence[tuple[int, ListNamespacesOp]],
results: list[Result],
cur: AsyncCursor[DictRow],
) -> None:
queries = self._get_batch_list_namespaces_queries(list_ops)
for (query, params), (idx, _) in zip(queries, list_ops):
await cur.execute(query, params)
rows = cast(list[dict], await cur.fetchall())
namespaces = [_decode_ns_bytes(row["truncated_prefix"]) for row in rows]
results[idx] = namespaces
@asynccontextmanager
async def _cursor(
self, *, pipeline: bool = False
) -> AsyncIterator[AsyncCursor[DictRow]]:
"""Create a database cursor as a context manager.
Args:
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.
"""
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:
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, 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.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()
@@ -370,5 +188,176 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
await cur.execute(sql)
await cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,))
async def _execute_batch(
self,
grouped_ops: dict,
results: list[Result],
conn: AsyncConnection[DictRow],
) -> None:
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]),
results,
cur,
)
if SearchOp in grouped_ops:
await self._batch_search_ops(
cast(Sequence[tuple[int, SearchOp]], grouped_ops[SearchOp]),
results,
cur,
)
if ListNamespacesOp in grouped_ops:
await self._batch_list_namespaces_ops(
cast(
Sequence[tuple[int, ListNamespacesOp]],
grouped_ops[ListNamespacesOp],
),
results,
cur,
)
if PutOp in grouped_ops:
await self._batch_put_ops(
cast(Sequence[tuple[int, PutOp]], grouped_ops[PutOp]),
cur,
)
async def _batch_get_ops(
self,
get_ops: Sequence[tuple[int, GetOp]],
results: list[Result],
cur: AsyncCursor[DictRow],
) -> None:
for query, params, namespace, items in self._get_batch_GET_ops_queries(get_ops):
await cur.execute(query, params)
rows = cast(list[Row], await cur.fetchall())
key_to_row = {row["key"]: row for row in rows}
for idx, key in items:
row = key_to_row.get(key)
if row:
results[idx] = _row_to_item(
namespace, row, loader=self._deserializer
)
else:
results[idx] = None
async def _batch_put_ops(
self,
put_ops: Sequence[tuple[int, PutOp]],
cur: AsyncCursor[DictRow],
) -> None:
queries, embedding_request = self._prepare_batch_PUT_queries(put_ops)
if embedding_request:
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__}."
)
query, txt_params = embedding_request
vectors = await self.embeddings.aembed_documents(
[param[-1] for param in txt_params]
)
queries.append(
(
query,
[
p
for (ns, k, pathname, _), vector in zip(txt_params, vectors)
for p in (ns, k, pathname, vector)
],
)
)
for query, params in queries:
await cur.execute(query, params)
async def _batch_search_ops(
self,
search_ops: Sequence[tuple[int, SearchOp]],
results: list[Result],
cur: AsyncCursor[DictRow],
) -> None:
queries, embedding_requests = self._prepare_batch_search_queries(search_ops)
if embedding_requests and self.embeddings:
vectors = await self.embeddings.aembed_documents(
[query for _, query in embedding_requests]
)
for (idx, _), vector in zip(embedding_requests, vectors):
queries[idx][1][0] = vector
for (idx, _), (query, params) in zip(search_ops, queries):
await cur.execute(query, params)
rows = cast(list[Row], await cur.fetchall())
items = [
_row_to_search_item(
_decode_ns_bytes(row["prefix"]), row, loader=self._deserializer
)
for row in rows
]
results[idx] = items
async def _batch_list_namespaces_ops(
self,
list_ops: Sequence[tuple[int, ListNamespacesOp]],
results: list[Result],
cur: AsyncCursor[DictRow],
) -> None:
queries = self._get_batch_list_namespaces_queries(list_ops)
for (query, params), (idx, _) in zip(queries, list_ops):
await cur.execute(query, params)
rows = cast(list[dict], await cur.fetchall())
namespaces = [_decode_ns_bytes(row["truncated_prefix"]) for row in rows]
results[idx] = namespaces
@asynccontextmanager
async def _cursor(
self, *, pipeline: bool = False
) -> AsyncIterator[AsyncCursor[DictRow]]:
"""Create a database cursor as a context manager.
Args:
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.
"""
async with _ainternal.get_connection(self.conn) as conn:
if self.pipe:
await self.pipe.sync()
# a connection in pipeline mode can be used concurrently
# in multiple threads/coroutines, but only one cursor can be
# used at a time
try:
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, 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.cursor(binary=True) as cur,
):
yield cur
@@ -356,11 +356,13 @@ class BasePostgresStore(Generic[C]):
k = op.key
for path, tokenized_path in paths:
for text in get_text_at_path(value, tokenized_path):
texts = get_text_at_path(value, tokenized_path)
for i, text in enumerate(texts):
pathname = f"{path}.{i}" if len(texts) > 1 else path
vector_values.append(
"(%s, %s, %s, %s, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)"
)
embedding_request_params.append((ns, k, path, text))
embedding_request_params.append((ns, k, pathname, text))
values_str = ",".join(values)
query = f"""
@@ -408,14 +410,13 @@ class BasePostgresStore(Generic[C]):
needs_vector_search = True
embedding_requests.append((idx, op.query))
_, score_expr = _get_distance_operator(self)
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"
@@ -424,14 +425,31 @@ class BasePostgresStore(Generic[C]):
else:
score_expr = score_expr % ("%s", vector_type)
vectors_per_doc_estimate = self.embedding_config[
"__estimated_num_vectors"
]
expanded_limit = (op.limit * vectors_per_doc_estimate * 2) + 1
# Direct query with DISTINCT ON to get best score per document
base_query = f"""
SELECT s.prefix, s.key, s.value, s.created_at, s.updated_at,
{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
with scored as (
SELECT DISTINCT ON (s.prefix, s.key)
s.prefix, s.key, s.value, s.created_at, s.updated_at,
{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
ORDER BY s.prefix, s.key, score DESC
LIMIT %s
)
SELECT * FROM scored
"""
params = [None, f"{_namespace_to_text(op.namespace_prefix)}%"]
params = [
None, # Vector placeholder
f"{_namespace_to_text(op.namespace_prefix)}%",
expanded_limit,
]
if op.filter:
filter_conditions = []
@@ -448,14 +466,17 @@ class BasePostgresStore(Generic[C]):
params.extend([key, json.dumps(value)])
if filter_conditions:
base_query += " AND " + " AND ".join(filter_conditions)
if needs_vector_search:
base_query += " WHERE " + " AND ".join(filter_conditions)
else:
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"
if needs_vector_search:
base_query += " ORDER BY score DESC"
else:
base_query += " ORDER BY updated_at DESC"
base_query += " LIMIT %s OFFSET %s"
params.extend([op.limit, op.offset])
queries.append((base_query, params))
@@ -554,18 +575,11 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
self.lock = threading.Lock()
self.embedding_config = embedding
if self.embedding_config:
self.embedding_config = self.embedding_config.copy()
self.embedding_config["__tokenized_fields"] = [
(p, tokenize_path(p)) if p != "__root__" else (p, p)
for p in (self.embedding_config.get("text_fields") or ["__root__"])
]
self.embeddings: Optional[Embeddings] = ensure_embeddings(
self.embedding_config.get("embed"),
aembed=self.embedding_config.get("aembed"),
self.embeddings, self.embedding_config = _ensure_embedding_config(
self.embedding_config
)
else:
self.embeddings = None
# TODO: Coerce embedding regular functions
@classmethod
@contextmanager
@@ -732,11 +746,15 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
[param[-1] for param in txt_params]
)
queries.extend(
[
(query, (ns, key, value, vector))
for (ns, key, value, _), vector in zip(txt_params, vectors)
]
queries.append(
(
query,
[
p
for (ns, k, pathname, _), vector in zip(txt_params, vectors)
for p in (ns, k, pathname, vector)
],
)
)
for query, params in queries:
@@ -825,9 +843,6 @@ class PostgresStore(BaseStore, BasePostgresStore[_pg_internal.Conn]):
cur.execute(sql)
cur.execute("INSERT INTO store_migrations (v) VALUES (%s)", (v,))
if self.pipe:
self.pipe.sync()
class Row(TypedDict):
key: str
@@ -934,17 +949,44 @@ def _decode_ns_bytes(namespace: Union[str, bytes, list]) -> tuple[str, ...]:
return tuple(namespace.split("."))
def _get_distance_operator(store: Any) -> tuple[str, str]:
def _get_distance_operator(store: Any) -> str:
"""Get the distance operator and score expression based on config."""
if not store.embedding_config:
return "<=>", "1 - (sv.embedding <=> %s::vector)"
raise ValueError(
"Embedding configuration is required for vector operations "
f"(for semantic search). "
f"Please provide an Embeddings when initializing the {store.__class__.__name__}."
)
config = cast(PostgresEmbeddingConfig, store.embedding_config)
distance_type = config.get("distance_type", "cosine")
if distance_type == "l2":
return "<->", "1 - (sv.embedding <-> %s::%s)"
return "1 - (sv.embedding <-> %s::%s)"
elif distance_type == "inner_product":
return "<#>", "-(sv.embedding <#> %s::%s)"
return "-(sv.embedding <#> %s::%s)"
else: # cosine
return "<=>", "1 - (sv.embedding <=> %s::%s)"
return "1 - (sv.embedding <=> %s::%s)"
def _ensure_embedding_config(
embedding_config: PostgresEmbeddingConfig,
) -> tuple[Optional["Embeddings"], PostgresEmbeddingConfig]:
embedding_config = embedding_config.copy()
tokenized: list[tuple[str, Union[Literal["__root__"], list[str]]]] = []
tot = 0
for p in embedding_config.get("text_fields") or ["__root__"]:
if p == "__root__":
tokenized.append((p, "__root__"))
tot += 1
else:
toks = tokenize_path(p)
tokenized.append((p, toks))
tot += len(toks)
embedding_config["__tokenized_fields"] = tokenized
embedding_config["__estimated_num_vectors"] = tot
embeddings = ensure_embeddings(
embedding_config.get("embed"),
aembed=embedding_config.get("aembed"),
)
return embeddings, embedding_config
+2 -1
View File
@@ -4,7 +4,8 @@ import pytest
from psycopg import AsyncConnection
from psycopg.errors import UndefinedTable
from psycopg.rows import DictRow, dict_row
from utils import CharacterEmbeddings # type: ignore
from langgraph.store.base._embed_test_utils import CharacterEmbeddings
DEFAULT_URI = "postgres://postgres:postgres@localhost:5441/postgres?sslmode=disable"
@@ -2,6 +2,8 @@
import sys
import uuid
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Any, Optional
import pytest
from conftest import (
@@ -187,6 +189,58 @@ async def test_batch_list_namespaces_ops(store: AsyncPostgresStore) -> None:
assert ("test", "namespace2") in results[0]
@asynccontextmanager
async def _create_vector_store(
index_type: str,
vector_type: str,
distance_type: str,
fake_embeddings: CharacterEmbeddings,
text_fields: Optional[list[str]] = None,
) -> 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
embedding_config = {
"dims": fake_embeddings.dims,
"embed": fake_embeddings,
"index_config": {
"kind": index_type,
"vector_type": vector_type,
},
"distance_type": distance_type,
"text_fields": text_fields,
}
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=embedding_config,
) 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}")
@pytest.fixture(
scope="function",
params=[
@@ -206,47 +260,11 @@ 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
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:
await conn.execute(f"CREATE DATABASE {database}")
try:
async with AsyncPostgresStore.from_conn_string(
conn_string,
embedding=embedding_config,
) 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 with _create_vector_store(
index_type, vector_type, distance_type, fake_embeddings
) as store:
yield store
async def test_vector_store_initialization(
@@ -398,3 +416,70 @@ async def test_vector_search_edge_cases(vector_store: AsyncPostgresStore) -> Non
results = await vector_store.asearch(("test",), query=special_query)
assert len(results) == 1
assert results[0].response_metadata["score"] < perfect_score
@pytest.mark.parametrize(
"index_type,vector_type,distance_type",
[
("ivfflat", "vector", "cosine"),
("hnsw", "vector", "cosine"),
("hnsw", "halfvec", "cosine"),
("hnsw", "halfvec", "inner_product"),
],
)
async def test_embed_with_path(
request: Any,
fake_embeddings: CharacterEmbeddings,
index_type: str,
vector_type: str,
distance_type: str,
) -> None:
"""Test vector search with specific text fields in Postgres store."""
async with _create_vector_store(
index_type,
vector_type,
distance_type,
fake_embeddings,
text_fields=["key0", "key1", "key3"],
) as store:
# This will have 2 vectors representing it
doc1 = {
# Omit key0 - check it doesn't raise an error
"key1": "xxx",
"key2": "yyy",
"key3": "zzz",
}
# This will have 3 vectors representing it
doc2 = {
"key0": "uuu",
"key1": "vvv",
"key2": "www",
"key3": "xxx",
}
await store.aput(("test",), "doc1", doc1)
await store.aput(("test",), "doc2", doc2)
# doc2.key3 and doc1.key1 both would have the highest score
results = await store.asearch(("test",), query="xxx")
assert len(results) == 2
assert results[0].key != results[1].key
ascore = results[0].response_metadata["score"]
bscore = results[1].response_metadata["score"]
assert ascore == pytest.approx(bscore, abs=1e-3)
results = await store.asearch(("test",), query="uuu")
assert len(results) == 2
assert results[0].key != results[1].key
assert results[0].key == "doc2"
assert (
results[0].response_metadata["score"]
> results[1].response_metadata["score"]
)
assert ascore == pytest.approx(results[0].response_metadata["score"], abs=1e-3)
# Un-indexed - will have low results for both. Not zero (because we're projecting)
# but less than the above.
results = await store.asearch(("test",), query="www")
assert len(results) == 2
assert results[0].response_metadata["score"] < ascore
assert results[1].response_metadata["score"] < ascore
+117 -16
View File
@@ -1,5 +1,7 @@
# type: ignore
from contextlib import contextmanager
from typing import Any, Optional
from uuid import uuid4
import pytest
@@ -348,21 +350,14 @@ class TestPostgresStore:
store.delete(namespace, key)
@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:
@contextmanager
def _create_vector_store(
index_type: str,
vector_type: str,
distance_type: str,
fake_embeddings: Embeddings,
text_fields: Optional[list[str]] = None,
) -> PostgresStore:
"""Create a store with vector search enabled."""
database = f"test_{uuid4().hex[:16]}"
uri_parts = DEFAULT_URI.split("/")
@@ -375,7 +370,6 @@ def vector_store(request, 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,
@@ -384,6 +378,7 @@ def vector_store(request, fake_embeddings: Embeddings) -> PostgresStore:
"vector_type": vector_type,
},
"distance_type": distance_type,
"text_fields": text_fields,
}
with Connection.connect(admin_conn_string, autocommit=True) as conn:
@@ -400,6 +395,32 @@ def vector_store(request, fake_embeddings: Embeddings) -> PostgresStore:
conn.execute(f"DROP DATABASE {database}")
@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."""
index_type, vector_type, distance_type = request.param
with _create_vector_store(
index_type, vector_type, distance_type, fake_embeddings
) as store:
yield store
def test_vector_store_initialization(
vector_store: PostgresStore, fake_embeddings: CharacterEmbeddings
) -> None:
@@ -533,3 +554,83 @@ def test_vector_search_edge_cases(vector_store: PostgresStore) -> None:
special_query = "test!@#$%^&*()"
results = vector_store.search(("test",), query=special_query)
assert len(results) == 1
@pytest.mark.parametrize(
"index_type,vector_type,distance_type",
[
("ivfflat", "vector", "cosine"),
("hnsw", "vector", "cosine"),
("hnsw", "halfvec", "cosine"),
("hnsw", "halfvec", "inner_product"),
],
)
def test_embed_with_path_sync(
request: Any,
fake_embeddings: CharacterEmbeddings,
index_type: str,
vector_type: str,
distance_type: str,
) -> None:
"""Test vector search with specific text fields in Postgres store."""
with _create_vector_store(
index_type,
vector_type,
distance_type,
fake_embeddings,
text_fields=["key0", "key1", "key3"],
) as store:
# This will have 2 vectors representing it
doc1 = {
# Omit key0 - check it doesn't raise an error
"key1": "xxx",
"key2": "yyy",
"key3": "zzz",
}
# This will have 3 vectors representing it
doc2 = {
"key0": "uuu",
"key1": "vvv",
"key2": "www",
"key3": "xxx",
}
store.put(("test",), "doc1", doc1)
store.put(("test",), "doc2", doc2)
# doc2.key3 and doc1.key1 both would have the highest score
results = store.search(("test",), query="xxx")
assert len(results) == 2
assert results[0].key != results[1].key
ascore = results[0].response_metadata["score"]
bscore = results[1].response_metadata["score"]
assert ascore == pytest.approx(bscore, abs=1e-3)
# ~Only match doc2
results = store.search(("test",), query="uuu")
assert len(results) == 2
assert results[0].key != results[1].key
assert results[0].key == "doc2"
assert (
results[0].response_metadata["score"]
> results[1].response_metadata["score"]
)
assert ascore == pytest.approx(results[0].response_metadata["score"], abs=1e-3)
# ~Only match doc1
results = store.search(("test",), query="zzz")
assert len(results) == 2
assert results[0].key != results[1].key
assert results[0].key == "doc1"
assert (
results[0].response_metadata["score"]
> results[1].response_metadata["score"]
)
assert ascore == pytest.approx(results[0].response_metadata["score"], abs=1e-3)
# Un-indexed - will have low results for both. Not zero (because we're projecting)
# but less than the above.
results = store.search(("test",), query="www")
assert len(results) == 2
assert results[0].key != results[1].key
assert results[0].response_metadata["score"] < ascore
assert results[1].response_metadata["score"] < ascore
-63
View File
@@ -1,63 +0,0 @@
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