fix(vector): harden backup restore parity gates

This commit is contained in:
2026-07-12 02:29:53 +02:00
parent e4db2ea5e1
commit 015c496bda
9 changed files with 300 additions and 36 deletions
+101 -21
View File
@@ -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]