Files
ThothII/harness/tests/l0/test_vector_adapter_parity.py
T

289 lines
12 KiB
Python

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, RestConfig
from tht.ports.vector import VectorRecord, VectorStoreError, VectorWriteRecord
from tht.vectorstore.rest_client import VectorRestClient, VectorRestError
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 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 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 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)
]
payload = sorted(
rows,
key=lambda row: (-row["similarity"], row["metadata"]["record_key"]),
)[: json["limit_count"]]
return Response(payload)
@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(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
@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")])
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]
def test_http_delete_generation_uses_exact_allowlisted_rpc_payload(monkeypatch):
calls = []
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda url, json, **kwargs: calls.append((url, json)) or Response({"deleted": 2}),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
assert client.delete_generation("evidence", "gen:" + "a" * 32, "default") == 2
assert calls == [("https://vectors.test/rpc/delete_vector_generation", {
"table_name": "evidence", "kind": "evidence", "generation": "gen:" + "a" * 32,
"workspace_id": "default",
})]
def test_http_delete_generation_legacy_404_fails_closed_without_body_leak(monkeypatch):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: Response({"message": "secret legacy endpoint detail"}, status=404),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
with pytest.raises(VectorRestError, match="delete_vector_generation RPC is unavailable") as error:
client.delete_generation("evidence", "gen:" + "a" * 32, "default")
assert "secret" not in str(error.value)
def test_http_list_evidence_generations_exact_rpc_and_legacy_fail_closed(monkeypatch):
calls = []
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda url, json, **kwargs: calls.append((url, json)) or Response([
{"generation": "gen:" + "a" * 32}
]),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
assert client.list_evidence_generations("evidence", "default") == ["gen:" + "a" * 32]
assert calls[0][0].endswith("/rpc/list_evidence_generations")
assert calls[0][1] == {"table_name": "evidence", "kind": "evidence", "workspace_id": "default"}
@pytest.mark.parametrize("generation", ["gen:a", "gen:" + "A" * 32, "gen:" + "a" * 33])
def test_http_generation_operations_reject_noncanonical_values(monkeypatch, generation):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: pytest.fail("invalid generation reached transport"),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
with pytest.raises(ValueError, match="canonical"):
client.delete_generation("evidence", generation, "default")
def test_http_inventory_rejects_malformed_rpc_output(monkeypatch):
monkeypatch.setattr(
"tht.vectorstore.rest_client.requests.post",
lambda *args, **kwargs: Response([{"generation": "gen:../escape"}]),
)
client = VectorRestClient(RestConfig(base_url="https://vectors.test", api_key="writer"))
with pytest.raises(VectorRestError, match="malformed"):
client.list_evidence_generations("evidence", "default")