140 lines
5.1 KiB
Python
140 lines
5.1 KiB
Python
import pytest
|
|
from sqlalchemy import create_engine
|
|
from testcontainers.postgres import PostgresContainer
|
|
|
|
from tht.config import DatabaseConfig
|
|
from tht.ports.vector import (
|
|
VectorReadUnavailable,
|
|
VectorRecord,
|
|
VectorStoreError,
|
|
VectorWriteRecord,
|
|
VectorWriteUnavailable,
|
|
)
|
|
|
|
|
|
def _record(content_hash: str, embedding: list[float], *, kind: str = "memory"):
|
|
return VectorWriteRecord(
|
|
record=VectorRecord(
|
|
id=f"record:{content_hash}",
|
|
kind=kind,
|
|
ref="session:test",
|
|
title=content_hash,
|
|
content=f"content {content_hash}",
|
|
metadata={"content_hash": content_hash},
|
|
),
|
|
embedding=embedding,
|
|
content_hash=content_hash,
|
|
)
|
|
|
|
|
|
@pytest.fixture(scope="module")
|
|
def vector_config():
|
|
with PostgresContainer("pgvector/pgvector:pg16") as pg:
|
|
host = pg.get_container_host_ip()
|
|
port = int(pg.get_exposed_port(5432))
|
|
config = DatabaseConfig(
|
|
host=host,
|
|
port=port,
|
|
database=pg.dbname,
|
|
schema="vectors",
|
|
user=pg.username,
|
|
password=pg.password,
|
|
)
|
|
engine = create_engine(pg.get_connection_url())
|
|
with engine.begin() as connection:
|
|
connection.exec_driver_sql("CREATE EXTENSION vector")
|
|
connection.exec_driver_sql("CREATE SCHEMA vectors")
|
|
for table in ("schema_records", "evidence", "memory"):
|
|
connection.exec_driver_sql(f"""
|
|
CREATE TABLE vectors.{table} (
|
|
id bigserial PRIMARY KEY,
|
|
record_key text UNIQUE NOT NULL,
|
|
kind text NOT NULL,
|
|
content_hash text NOT NULL,
|
|
metadata jsonb NOT NULL,
|
|
embedding vector(2) NOT NULL,
|
|
indexed_at timestamptz NOT NULL DEFAULT now()
|
|
)
|
|
""")
|
|
engine.dispose()
|
|
yield config
|
|
|
|
|
|
@pytest.fixture
|
|
def store(vector_config):
|
|
from tht.adapters.vector.pgvector import PgVectorStore
|
|
|
|
store = PgVectorStore(vector_config, vector_config, expected_dimension=2)
|
|
store.upsert("memory", [_record("reset", [0.0, 1.0])])
|
|
yield store
|
|
|
|
|
|
def test_pgvector_round_trip_hash_and_upsert(store):
|
|
assert store.upsert("memory", [_record("a", [1.0, 0.0])]) == 1
|
|
assert store.existing_hashes("memory", ["memory"])["record:a"] == "a"
|
|
|
|
hits = store.search(["memory"], [1.0, 0.0], limit=5, kinds=["memory"])
|
|
assert hits[0].metadata["content_hash"] == "a"
|
|
assert hits[0].id == "record:a"
|
|
|
|
assert store.upsert("memory", [_record("a", [0.8, 0.2])]) == 1
|
|
assert store.search(["memory"], [0.8, 0.2], limit=1)[0].id == "record:a"
|
|
|
|
|
|
def test_pgvector_search_filters_kinds_before_limit(store):
|
|
store.upsert("memory", [_record("solved", [1.0, 0.0], kind="solved_question")])
|
|
hits = store.search("memory".split(), [1.0, 0.0], limit=1, kinds=["memory"])
|
|
assert len(hits) == 1
|
|
assert hits[0].kind == "memory"
|
|
|
|
|
|
@pytest.mark.parametrize("limit", [True, False, 1.0, 0, -1])
|
|
def test_pgvector_search_requires_strict_positive_limit(store, limit):
|
|
with pytest.raises(ValueError, match="positive integer"):
|
|
store.search(["memory"], [1.0, 0.0], limit=limit)
|
|
|
|
|
|
def test_pgvector_allowlists_collections(store):
|
|
with pytest.raises(VectorStoreError, match="Collection not allowed"):
|
|
store.search(["memory; DROP SCHEMA vectors"], [1.0, 0.0], limit=1)
|
|
with pytest.raises(VectorStoreError, match="Collection not allowed"):
|
|
store.upsert("unknown", [])
|
|
|
|
|
|
def test_pgvector_rejects_kinds_not_belonging_to_collection(store):
|
|
with pytest.raises(VectorStoreError, match="Kind not allowed"):
|
|
store.existing_hashes("evidence", ["memory"])
|
|
with pytest.raises(VectorStoreError, match="Kind not allowed"):
|
|
store.upsert("evidence", [_record("wrong", [1.0, 0.0])])
|
|
|
|
|
|
def test_pgvector_separates_read_and_write_credentials(vector_config):
|
|
from tht.adapters.vector.pgvector import PgVectorStore
|
|
|
|
reader = PgVectorStore(vector_config, expected_dimension=2)
|
|
assert reader.capabilities.search is True
|
|
assert reader.capabilities.upsert is False
|
|
with pytest.raises(VectorWriteUnavailable):
|
|
reader.upsert("memory", [])
|
|
|
|
writer = PgVectorStore(None, vector_config, expected_dimension=2)
|
|
assert writer.capabilities.search is False
|
|
assert writer.capabilities.upsert is True
|
|
with pytest.raises(VectorReadUnavailable):
|
|
writer.search(["memory"], [1.0, 0.0], limit=1)
|
|
|
|
|
|
def test_pgvector_health_reports_dimension_and_each_connection(vector_config):
|
|
from tht.adapters.vector.pgvector import PgVectorStore
|
|
|
|
health = PgVectorStore(vector_config, vector_config, expected_dimension=2).health()
|
|
assert health.ok is True
|
|
assert health.read_reachable is True
|
|
assert health.write_reachable is True
|
|
assert health.observed_dimensions == (2,)
|
|
assert health.dimension_compatible is True
|
|
|
|
mismatch = PgVectorStore(vector_config, None, expected_dimension=3).health()
|
|
assert mismatch.ok is False
|
|
assert mismatch.dimension_compatible is False
|