docs(vector): add local backup restore and parity gate
This commit is contained in:
@@ -0,0 +1,148 @@
|
||||
import math
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine
|
||||
from testcontainers.postgres import PostgresContainer
|
||||
|
||||
from tht.adapters.vector.pgvector import PgVectorStore
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
from tht.config import DatabaseConfig
|
||||
from tht.ports.vector import VectorRecord, VectorStoreError, VectorWriteRecord
|
||||
|
||||
|
||||
def _write(record_id, kind, embedding, content_hash):
|
||||
return VectorWriteRecord(
|
||||
VectorRecord(
|
||||
id=record_id,
|
||||
kind=kind,
|
||||
ref="fixture",
|
||||
title=record_id,
|
||||
content=f"content {record_id}",
|
||||
metadata={"fixture": True},
|
||||
),
|
||||
embedding,
|
||||
content_hash,
|
||||
)
|
||||
|
||||
|
||||
FIXTURE = [
|
||||
_write("memory:a", "memory", [1.0, 0.0], "hash-a"),
|
||||
_write("memory:b", "memory", [1.0, 0.0], "hash-b"),
|
||||
_write("solved:a", "solved_question", [0.8, 0.2], "hash-solved"),
|
||||
]
|
||||
|
||||
|
||||
class FixtureHttpClient:
|
||||
def __init__(self):
|
||||
self.rows = {}
|
||||
|
||||
def list_tables(self):
|
||||
return [{"table_name": "memory", "vector_dimensions": 2}]
|
||||
|
||||
def upsert_records(self, table_name, rows):
|
||||
for row in rows:
|
||||
self.rows[(table_name, row["record_key"])] = row
|
||||
return len(rows)
|
||||
|
||||
def existing_hashes(self, table_name, kinds):
|
||||
return {
|
||||
row["record_key"]: row["content_hash"]
|
||||
for (table, _), row in self.rows.items()
|
||||
if table == table_name and row["kind"] in kinds
|
||||
}
|
||||
|
||||
def search_similar(self, table_name, embedding, limit, kinds=None):
|
||||
def similarity(row):
|
||||
left, right = row["embedding"], embedding
|
||||
return sum(a * b for a, b in zip(left, right)) / (
|
||||
math.sqrt(sum(a * a for a in left))
|
||||
* math.sqrt(sum(b * b for b in right))
|
||||
)
|
||||
|
||||
rows = [
|
||||
{"metadata": row["metadata"], "similarity": similarity(row)}
|
||||
for (table, _), row in self.rows.items()
|
||||
if table == table_name and (not kinds or row["kind"] in kinds)
|
||||
]
|
||||
return sorted(
|
||||
rows,
|
||||
key=lambda row: (-row["similarity"], row["metadata"]["record_key"]),
|
||||
)[:limit]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def direct_store():
|
||||
with PostgresContainer("pgvector/pgvector:pg16") as postgres:
|
||||
config = DatabaseConfig(
|
||||
host=postgres.get_container_host_ip(),
|
||||
port=int(postgres.get_exposed_port(5432)),
|
||||
database=postgres.dbname,
|
||||
schema="vectors",
|
||||
user=postgres.username,
|
||||
password=postgres.password,
|
||||
)
|
||||
engine = create_engine(postgres.get_connection_url())
|
||||
with engine.begin() as connection:
|
||||
connection.exec_driver_sql("CREATE SCHEMA vectors")
|
||||
connection.exec_driver_sql("CREATE EXTENSION vector WITH SCHEMA vectors")
|
||||
connection.exec_driver_sql(
|
||||
"CREATE TABLE vectors.memory ("
|
||||
"id bigserial PRIMARY KEY, record_key text UNIQUE NOT NULL, "
|
||||
"kind text NOT NULL, content_hash text NOT NULL, metadata jsonb NOT NULL, "
|
||||
"embedding vectors.vector(2) NOT NULL, indexed_at timestamptz NOT NULL "
|
||||
"DEFAULT now())"
|
||||
)
|
||||
engine.dispose()
|
||||
reader, writer = config, config
|
||||
store = PgVectorStore(reader, writer, expected_dimension=2)
|
||||
store.upsert("memory", FIXTURE)
|
||||
yield store
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def http_store():
|
||||
client = FixtureHttpClient()
|
||||
store = ThothHttpVectorStore(client, client, expected_dimension=2)
|
||||
store.upsert("memory", FIXTURE)
|
||||
return store
|
||||
|
||||
|
||||
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
|
||||
def test_kind_filtered_search_has_identical_order(request, store_fixture):
|
||||
store = request.getfixturevalue(store_fixture)
|
||||
hits = store.search(["memory"], [1.0, 0.0], limit=3, kinds=["memory"])
|
||||
assert [(hit.id, hit.kind, round(hit.similarity, 6)) for hit in hits] == [
|
||||
("memory:a", "memory", 1.0),
|
||||
("memory:b", "memory", 1.0),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
|
||||
def test_hash_and_upsert_parity(request, store_fixture):
|
||||
store = request.getfixturevalue(store_fixture)
|
||||
assert store.existing_hashes("memory", ["memory"]) == {
|
||||
"memory:a": "hash-a",
|
||||
"memory:b": "hash-b",
|
||||
}
|
||||
replacement = _write("memory:a", "memory", [0.0, 1.0], "hash-a-2")
|
||||
assert store.upsert("memory", [replacement]) == 1
|
||||
assert store.existing_hashes("memory", ["memory"])["memory:a"] == "hash-a-2"
|
||||
assert store.search(["memory"], [0.0, 1.0], limit=1, kinds=["memory"])[0].id == "memory:a"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
|
||||
def test_validation_error_parity(request, store_fixture):
|
||||
store = request.getfixturevalue(store_fixture)
|
||||
with pytest.raises(VectorStoreError, match="Collection not allowed"):
|
||||
store.search(["not_allowed"], [1.0, 0.0], limit=1)
|
||||
with pytest.raises(VectorStoreError, match="Kind not allowed"):
|
||||
store.search(["memory"], [1.0, 0.0], limit=1, kinds=["not_allowed"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("store_fixture", ["direct_store", "http_store"])
|
||||
def test_dimension_error_parity(request, store_fixture):
|
||||
store = request.getfixturevalue(store_fixture)
|
||||
with pytest.raises(VectorStoreError, match="Query embedding dimension"):
|
||||
store.search(["memory"], [1.0], limit=1)
|
||||
with pytest.raises(VectorStoreError, match="Embedding dimension"):
|
||||
store.upsert("memory", [_write("bad", "memory", [1.0], "bad")])
|
||||
@@ -257,7 +257,7 @@ class PgVectorStore:
|
||||
where = sql.SQL(" WHERE kind = ANY(%s)") if collection_kinds else sql.SQL("")
|
||||
query = sql.SQL(
|
||||
"SELECT metadata, 1 - (embedding {} %s::{}) AS similarity "
|
||||
"FROM {}{} ORDER BY embedding {} %s::{} LIMIT %s"
|
||||
"FROM {}{} ORDER BY embedding {} %s::{}, record_key LIMIT %s"
|
||||
).format(
|
||||
_cosine_operator(self._schema),
|
||||
_vector_type(self._schema),
|
||||
@@ -274,7 +274,7 @@ class PgVectorStore:
|
||||
hits.extend(hit_from_metadata(row[1], row[0]) for row in cursor.fetchall())
|
||||
finally:
|
||||
raw.close()
|
||||
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[:limit]
|
||||
return sorted(hits, key=lambda hit: (-hit.similarity, hit.id))[:limit]
|
||||
|
||||
def _require_writer(self) -> Engine:
|
||||
if self._writer is None:
|
||||
|
||||
@@ -5,16 +5,22 @@ from tht.ports.vector import (
|
||||
VectorHealth,
|
||||
VectorHit,
|
||||
VectorReadUnavailable,
|
||||
VectorStoreError,
|
||||
VectorWriteRecord,
|
||||
VectorWriteUnavailable,
|
||||
require_positive_limit,
|
||||
)
|
||||
from tht.vectorstore.rest_client import VectorRestClient
|
||||
from tht.vectorstore.store import hit_from_metadata
|
||||
from tht.adapters.vector.pgvector import (
|
||||
_collection,
|
||||
_validate_collection_kinds,
|
||||
_validate_known_kinds,
|
||||
)
|
||||
|
||||
|
||||
def _merge(hits: list[VectorHit], limit: int) -> list[VectorHit]:
|
||||
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[:limit]
|
||||
return sorted(hits, key=lambda hit: (-hit.similarity, hit.id))[:limit]
|
||||
|
||||
|
||||
class ThothHttpVectorStore:
|
||||
@@ -89,8 +95,13 @@ class ThothHttpVectorStore:
|
||||
require_positive_limit(limit)
|
||||
if self._reader is None:
|
||||
raise VectorReadUnavailable("Vector reader credential is not configured")
|
||||
if self._expected_dimension is not None and len(embedding) != self._expected_dimension:
|
||||
raise VectorStoreError("Query embedding dimension does not match configured dimension")
|
||||
if kinds:
|
||||
_validate_known_kinds(kinds)
|
||||
hits: list[VectorHit] = []
|
||||
for collection in collections:
|
||||
_collection("vectors", collection)
|
||||
rows = self._reader.search_similar(collection, embedding, limit, kinds=kinds)
|
||||
hits.extend(
|
||||
hit_from_metadata(row.get("similarity", 0.0), row.get("metadata"))
|
||||
@@ -107,10 +118,20 @@ class ThothHttpVectorStore:
|
||||
return self._writer
|
||||
|
||||
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
|
||||
_collection("vectors", collection)
|
||||
_validate_collection_kinds(collection, kinds)
|
||||
return self._require_writer().existing_hashes(collection, kinds)
|
||||
|
||||
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
|
||||
writer = self._require_writer()
|
||||
_collection("vectors", collection)
|
||||
for record in records:
|
||||
_validate_collection_kinds(collection, [record.record.kind])
|
||||
if (
|
||||
self._expected_dimension is not None
|
||||
and len(record.embedding) != self._expected_dimension
|
||||
):
|
||||
raise VectorStoreError("Embedding dimension does not match configured dimension")
|
||||
rows = [self._row(record) for record in records]
|
||||
return writer.upsert_records(collection, rows)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user