mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-11 04:07:52 +02:00
merge
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 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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -4,11 +4,13 @@
|
||||
# TESTING AND COVERAGE
|
||||
######################
|
||||
|
||||
TEST ?= .
|
||||
|
||||
test:
|
||||
poetry run pytest tests
|
||||
poetry run pytest $(TEST)
|
||||
|
||||
test_watch:
|
||||
poetry run ptw .
|
||||
poetry run ptw $(TEST)
|
||||
|
||||
######################
|
||||
# LINTING AND FORMATTING
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
|
||||
import math
|
||||
import random
|
||||
from collections import Counter
|
||||
from typing import Any, Optional
|
||||
from collections import Counter, defaultdict
|
||||
from typing import Any
|
||||
|
||||
from langchain_core.embeddings import Embeddings
|
||||
|
||||
@@ -14,36 +14,28 @@ class CharacterEmbeddings(Embeddings):
|
||||
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))
|
||||
# Create projection vector for each character lazily
|
||||
self._char_projections: defaultdict[str, list[float]] = defaultdict(
|
||||
lambda: [
|
||||
self._rng.gauss(0, 1 / math.sqrt(self.dims)) for _ in range(self.dims)
|
||||
]
|
||||
)
|
||||
|
||||
def _embed_one(self, text: str) -> list[float]:
|
||||
"""Embed a single text."""
|
||||
counts = Counter(text)
|
||||
char_vec = [0.0] * len(self._char_to_idx)
|
||||
total = sum(counts.values())
|
||||
|
||||
if total == 0:
|
||||
return [0.0] * self.dims
|
||||
|
||||
embedding = [0.0] * self.dims
|
||||
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)
|
||||
]
|
||||
weight = count / total
|
||||
char_proj = self._char_projections[char]
|
||||
for i, proj in enumerate(char_proj):
|
||||
embedding[i] += weight * proj
|
||||
|
||||
norm = math.sqrt(sum(x * x for x in embedding))
|
||||
if norm > 0:
|
||||
@@ -53,12 +45,10 @@ class CharacterEmbeddings(Embeddings):
|
||||
|
||||
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:
|
||||
|
||||
@@ -199,7 +199,7 @@ class InMemoryStore(BaseStore):
|
||||
with cf.ThreadPoolExecutor() as executor:
|
||||
futures = {
|
||||
q: executor.submit(self.embeddings.embed_query, q)
|
||||
for q in queries
|
||||
for q in list(queries)
|
||||
}
|
||||
for query, future in futures.items():
|
||||
queryinmem_store[query] = future.result()
|
||||
@@ -215,7 +215,7 @@ class InMemoryStore(BaseStore):
|
||||
queries = {op.query for (op, _) in search_ops.values() if op.query}
|
||||
|
||||
if queries:
|
||||
coros = [self.embeddings.aembed_query(q) for q in queries]
|
||||
coros = [self.embeddings.aembed_query(q) for q in list(queries)]
|
||||
results = await asyncio.gather(*coros)
|
||||
queryinmem_store = dict(zip(queries, results))
|
||||
|
||||
@@ -404,15 +404,15 @@ def _cosine_similarity(X: list[float], Y: list[list[float]]) -> list[float]:
|
||||
if _check_numpy():
|
||||
import numpy as np # type: ignore
|
||||
|
||||
X = np.array(X) if not isinstance(X, np.ndarray) else X
|
||||
Y = np.array(Y) if not isinstance(Y, np.ndarray) else Y
|
||||
X_norm = np.linalg.norm(X)
|
||||
Y_norm = np.linalg.norm(Y, axis=1)
|
||||
X_arr = np.array(X) if not isinstance(X, np.ndarray) else X
|
||||
Y_arr = np.array(Y) if not isinstance(Y, np.ndarray) else Y
|
||||
X_norm = np.linalg.norm(X_arr)
|
||||
Y_norm = np.linalg.norm(Y_arr, axis=1)
|
||||
|
||||
# Avoid division by zero
|
||||
mask = Y_norm != 0
|
||||
similarities = np.zeros_like(Y_norm)
|
||||
similarities[mask] = np.dot(Y[mask], X) / (Y_norm[mask] * X_norm)
|
||||
similarities[mask] = np.dot(Y_arr[mask], X_arr) / (Y_norm[mask] * X_norm)
|
||||
return similarities.tolist()
|
||||
|
||||
similarities = []
|
||||
|
||||
@@ -1,11 +1,20 @@
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Iterable
|
||||
|
||||
import pytest
|
||||
from pytest_mock import MockerFixture
|
||||
|
||||
from langgraph.store.base import GetOp, InvalidNamespaceError, Item, Op, PutOp, Result
|
||||
from langgraph.store.base import (
|
||||
GetOp,
|
||||
InvalidNamespaceError,
|
||||
Item,
|
||||
Op,
|
||||
PutOp,
|
||||
Result,
|
||||
get_text_at_path,
|
||||
)
|
||||
from langgraph.store.base._embed_test_utils import CharacterEmbeddings
|
||||
from langgraph.store.base.batch import AsyncBatchedBaseStore
|
||||
from langgraph.store.memory import InMemoryStore
|
||||
@@ -23,6 +32,74 @@ class MockAsyncBatchedStore(AsyncBatchedBaseStore):
|
||||
return self._store.batch(ops)
|
||||
|
||||
|
||||
def test_get_text_at_path() -> None:
|
||||
nested_data = {
|
||||
"name": "test",
|
||||
"info": {
|
||||
"age": 25,
|
||||
"tags": ["a", "b", "c"],
|
||||
"metadata": {"created": "2024-01-01", "updated": "2024-01-02"},
|
||||
},
|
||||
"items": [
|
||||
{"id": 1, "value": "first", "tags": ["x", "y"]},
|
||||
{"id": 2, "value": "second", "tags": ["y", "z"]},
|
||||
{"id": 3, "value": "third", "tags": ["z", "w"]},
|
||||
],
|
||||
"empty": None,
|
||||
"zeros": [0, 0.0, "0"],
|
||||
"empty_list": [],
|
||||
"empty_dict": {},
|
||||
}
|
||||
|
||||
assert get_text_at_path(nested_data, "__root__") == [
|
||||
json.dumps(nested_data, sort_keys=True)
|
||||
]
|
||||
|
||||
assert get_text_at_path(nested_data, "name") == ["test"]
|
||||
assert get_text_at_path(nested_data, "info.age") == ["25"]
|
||||
|
||||
assert get_text_at_path(nested_data, "info.metadata.created") == ["2024-01-01"]
|
||||
|
||||
assert get_text_at_path(nested_data, "items[0].value") == ["first"]
|
||||
assert get_text_at_path(nested_data, "items[-1].value") == ["third"]
|
||||
assert get_text_at_path(nested_data, "items[1].tags[0]") == ["y"]
|
||||
|
||||
values = get_text_at_path(nested_data, "items[*].value")
|
||||
assert set(values) == {"first", "second", "third"}
|
||||
|
||||
metadata_dates = get_text_at_path(nested_data, "info.metadata.*")
|
||||
assert set(metadata_dates) == {"2024-01-01", "2024-01-02"}
|
||||
name_and_age = get_text_at_path(nested_data, "{name,info.age}")
|
||||
assert set(name_and_age) == {"test", "25"}
|
||||
|
||||
item_fields = get_text_at_path(nested_data, "items[*].{id,value}")
|
||||
assert set(item_fields) == {"1", "2", "3", "first", "second", "third"}
|
||||
|
||||
all_tags = get_text_at_path(nested_data, "items[*].tags[*]")
|
||||
assert set(all_tags) == {"x", "y", "z", "w"}
|
||||
|
||||
assert get_text_at_path(None, "any.path") == []
|
||||
assert get_text_at_path({}, "any.path") == []
|
||||
assert get_text_at_path(nested_data, "") == [
|
||||
json.dumps(nested_data, sort_keys=True)
|
||||
]
|
||||
assert get_text_at_path(nested_data, "nonexistent") == []
|
||||
assert get_text_at_path(nested_data, "items[99].value") == []
|
||||
assert get_text_at_path(nested_data, "items[*].nonexistent") == []
|
||||
|
||||
assert get_text_at_path(nested_data, "empty") == []
|
||||
assert get_text_at_path(nested_data, "empty_list") == ["[]"]
|
||||
assert get_text_at_path(nested_data, "empty_dict") == ["{}"]
|
||||
|
||||
zeros = get_text_at_path(nested_data, "zeros[*]")
|
||||
assert set(zeros) == {"0", "0.0"}
|
||||
|
||||
assert get_text_at_path(nested_data, "items[].value") == []
|
||||
assert get_text_at_path(nested_data, "items[abc].value") == []
|
||||
assert get_text_at_path(nested_data, "{unclosed") == []
|
||||
assert get_text_at_path(nested_data, "nested[{invalid}]") == []
|
||||
|
||||
|
||||
async def test_async_batch_store(mocker: MockerFixture) -> None:
|
||||
abatch = mocker.stub()
|
||||
|
||||
@@ -804,3 +881,53 @@ async def test_async_vector_search_edge_cases(
|
||||
special_query = "test!@#$%^&*()"
|
||||
results = await store.asearch(("test",), query=special_query)
|
||||
assert len(results) == 1
|
||||
|
||||
|
||||
async def test_embed_with_path(fake_embeddings: CharacterEmbeddings) -> None:
|
||||
# Basi
|
||||
store = InMemoryStore(
|
||||
embedding_config={
|
||||
"dims": fake_embeddings.dims,
|
||||
"embed": fake_embeddings,
|
||||
# Key 2 isn't included. Don't index it.
|
||||
"text_fields": ["key0", "key1", "key3"],
|
||||
}
|
||||
)
|
||||
# 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 == bscore
|
||||
|
||||
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-5)
|
||||
|
||||
# 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
|
||||
|
||||
Reference in New Issue
Block a user