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")