fix(vector): harden backup restore parity gates
This commit is contained in:
@@ -6,8 +6,9 @@ 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.config import DatabaseConfig, RestConfig
|
||||
from tht.ports.vector import VectorRecord, VectorStoreError, VectorWriteRecord
|
||||
from tht.vectorstore.rest_client import VectorRestClient
|
||||
|
||||
|
||||
def _write(record_id, kind, embedding, content_hash):
|
||||
@@ -32,26 +33,46 @@ FIXTURE = [
|
||||
]
|
||||
|
||||
|
||||
class FixtureHttpClient:
|
||||
class Response:
|
||||
def __init__(self, payload=None, status=200):
|
||||
self.status_code = status
|
||||
self.payload = payload
|
||||
self.text = "" if payload is None else "json"
|
||||
|
||||
@property
|
||||
def ok(self):
|
||||
return self.status_code < 400
|
||||
|
||||
def json(self):
|
||||
return self.payload
|
||||
|
||||
|
||||
class FixtureHttpTransport:
|
||||
def __init__(self):
|
||||
self.rows = {}
|
||||
self.calls = []
|
||||
|
||||
def list_tables(self):
|
||||
return [{"table_name": "memory", "vector_dimensions": 2}]
|
||||
def post(self, url, json, headers, **kwargs):
|
||||
assert headers == {"X-API-Key": "parity-key"}
|
||||
self.calls.append((url.rsplit("/", 1)[-1], json))
|
||||
function = self.calls[-1][0]
|
||||
if function == "list_tables":
|
||||
return Response([{"table_name": "memory", "vector_dimensions": 2}])
|
||||
if function == "upsert_vector_records":
|
||||
for row in json["rows"]:
|
||||
self.rows[(json["table_name"], row["record_key"])] = row
|
||||
return Response({"upserted": len(json["rows"])})
|
||||
if function == "existing_vector_hashes":
|
||||
return Response([
|
||||
{"record_key": row["record_key"], "content_hash": row["content_hash"]}
|
||||
for (table, _), row in self.rows.items()
|
||||
if table == json["table_name"] and row["kind"] in json["kinds"]
|
||||
])
|
||||
assert function == "search_similar"
|
||||
table_name = json["table_name"]
|
||||
embedding = json["query_embedding"]
|
||||
kinds = json.get("kinds")
|
||||
|
||||
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)) / (
|
||||
@@ -64,10 +85,11 @@ class FixtureHttpClient:
|
||||
for (table, _), row in self.rows.items()
|
||||
if table == table_name and (not kinds or row["kind"] in kinds)
|
||||
]
|
||||
return sorted(
|
||||
payload = sorted(
|
||||
rows,
|
||||
key=lambda row: (-row["similarity"], row["metadata"]["record_key"]),
|
||||
)[:limit]
|
||||
)[: json["limit_count"]]
|
||||
return Response(payload)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -100,10 +122,13 @@ def direct_store():
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def http_store():
|
||||
client = FixtureHttpClient()
|
||||
def http_store(monkeypatch):
|
||||
transport = FixtureHttpTransport()
|
||||
monkeypatch.setattr("tht.vectorstore.rest_client.requests.post", transport.post)
|
||||
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
|
||||
store = ThothHttpVectorStore(client, client, expected_dimension=2)
|
||||
store.upsert("memory", FIXTURE)
|
||||
store.transport = transport
|
||||
return store
|
||||
|
||||
|
||||
@@ -146,3 +171,58 @@ def test_dimension_error_parity(request, store_fixture):
|
||||
store.search(["memory"], [1.0], limit=1)
|
||||
with pytest.raises(VectorStoreError, match="Embedding dimension"):
|
||||
store.upsert("memory", [_write("bad", "memory", [1.0], "bad")])
|
||||
|
||||
|
||||
def test_http_parity_exercises_rpc_kinds_payload(http_store):
|
||||
http_store.search(["memory"], [1.0, 0.0], limit=2, kinds=["memory"])
|
||||
search_calls = [payload for function, payload in http_store.transport.calls if function == "search_similar"]
|
||||
assert search_calls[-1] == {
|
||||
"query_embedding": [1.0, 0.0],
|
||||
"limit_count": 2,
|
||||
"table_name": "memory",
|
||||
"kinds": ["memory"],
|
||||
}
|
||||
|
||||
|
||||
def test_http_adapter_maps_transport_error(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"tht.vectorstore.rest_client.requests.post",
|
||||
lambda *args, **kwargs: Response({"message": "server broke"}, status=500),
|
||||
)
|
||||
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
|
||||
store = ThothHttpVectorStore(client, client, expected_dimension=2)
|
||||
with pytest.raises(VectorStoreError, match="HTTP 500"):
|
||||
store.search(["memory"], [1.0, 0.0], limit=1, kinds=["memory"])
|
||||
|
||||
|
||||
def test_http_adapter_tolerates_malformed_metadata(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"tht.vectorstore.rest_client.requests.post",
|
||||
lambda *args, **kwargs: Response([{"similarity": 0.5, "metadata": None}]),
|
||||
)
|
||||
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
|
||||
hit = ThothHttpVectorStore(client, None, expected_dimension=2).search(
|
||||
["memory"], [1.0, 0.0], limit=1
|
||||
)[0]
|
||||
assert (hit.id, hit.kind, hit.metadata) == ("", "", {})
|
||||
|
||||
|
||||
def test_http_adapter_legacy_fallback_preserves_kind_semantics(monkeypatch):
|
||||
calls = []
|
||||
|
||||
def post(url, json, **kwargs):
|
||||
calls.append(json)
|
||||
if "kinds" in json:
|
||||
return Response({"message": "function not found"}, status=404)
|
||||
return Response([
|
||||
{"similarity": 1.0, "metadata": {"record_key": "wrong", "kind": "solved_question"}},
|
||||
{"similarity": 0.9, "metadata": {"record_key": "right", "kind": "memory"}},
|
||||
])
|
||||
|
||||
monkeypatch.setattr("tht.vectorstore.rest_client.requests.post", post)
|
||||
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="parity-key"))
|
||||
hits = ThothHttpVectorStore(client, None, expected_dimension=2).search(
|
||||
["memory"], [1.0, 0.0], limit=2, kinds=["memory"]
|
||||
)
|
||||
assert [hit.id for hit in hits] == ["right"]
|
||||
assert "kinds" in calls[0] and "kinds" not in calls[1]
|
||||
|
||||
@@ -10,7 +10,7 @@ from tht.ports.vector import (
|
||||
VectorWriteUnavailable,
|
||||
require_positive_limit,
|
||||
)
|
||||
from tht.vectorstore.rest_client import VectorRestClient
|
||||
from tht.vectorstore.rest_client import VectorRestClient, VectorRestError
|
||||
from tht.vectorstore.store import hit_from_metadata
|
||||
from tht.adapters.vector.pgvector import (
|
||||
_collection,
|
||||
@@ -102,7 +102,10 @@ class ThothHttpVectorStore:
|
||||
hits: list[VectorHit] = []
|
||||
for collection in collections:
|
||||
_collection("vectors", collection)
|
||||
rows = self._reader.search_similar(collection, embedding, limit, kinds=kinds)
|
||||
try:
|
||||
rows = self._reader.search_similar(collection, embedding, limit, kinds=kinds)
|
||||
except VectorRestError as exc:
|
||||
raise VectorStoreError(str(exc)) from exc
|
||||
hits.extend(
|
||||
hit_from_metadata(row.get("similarity", 0.0), row.get("metadata"))
|
||||
for row in rows
|
||||
@@ -120,7 +123,10 @@ class ThothHttpVectorStore:
|
||||
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)
|
||||
try:
|
||||
return self._require_writer().existing_hashes(collection, kinds)
|
||||
except VectorRestError as exc:
|
||||
raise VectorStoreError(str(exc)) from exc
|
||||
|
||||
def upsert(self, collection: str, records: list[VectorWriteRecord]) -> int:
|
||||
writer = self._require_writer()
|
||||
@@ -133,7 +139,10 @@ class ThothHttpVectorStore:
|
||||
):
|
||||
raise VectorStoreError("Embedding dimension does not match configured dimension")
|
||||
rows = [self._row(record) for record in records]
|
||||
return writer.upsert_records(collection, rows)
|
||||
try:
|
||||
return writer.upsert_records(collection, rows)
|
||||
except VectorRestError as exc:
|
||||
raise VectorStoreError(str(exc)) from exc
|
||||
|
||||
@staticmethod
|
||||
def _row(write_record: VectorWriteRecord) -> dict:
|
||||
|
||||
Reference in New Issue
Block a user