289 lines
12 KiB
Python
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")
|