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

206 lines
7.5 KiB
Python

import json
from pathlib import Path
import pytest
from sqlalchemy import create_engine, text
from sqlalchemy.exc import ProgrammingError
from testcontainers.postgres import PostgresContainer
from typer.testing import CliRunner
from tht.cli import app
from tht.config import DatabaseConfig
from tht.ports.vector import VectorRecord, VectorWriteRecord
@pytest.fixture(scope="module")
def database_url():
with PostgresContainer("pgvector/pgvector:pg16") as postgres:
yield postgres.get_connection_url()
def test_migrations_are_clean_and_idempotent(database_url):
from tht.cli.vector_migrate_cmd import migrate, migration_status
before = migration_status(database_url)
assert [item.version for item in before.pending] == ["001", "002", "003"]
migrate(database_url)
migrate(database_url)
status = migration_status(database_url)
assert status.pending == ()
assert status.drifted == ()
assert [item.version for item in status.applied] == ["001", "002", "003"]
def test_schema_matches_direct_adapter_contract(database_url):
engine = create_engine(database_url)
with engine.connect() as connection:
rows = connection.execute(
text(
"SELECT table_name, column_name, data_type, udt_name "
"FROM information_schema.columns WHERE table_schema = 'vectors' "
"ORDER BY table_name, ordinal_position"
)
).all()
vector_types = connection.execute(
text(
"SELECT c.relname, format_type(a.atttypid, a.atttypmod) "
"FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace "
"JOIN pg_attribute a ON a.attrelid = c.oid AND a.attname = 'embedding' "
"WHERE n.nspname = 'vectors' ORDER BY c.relname"
)
).all()
engine.dispose()
tables = {row.table_name for row in rows}
assert tables == {"evidence", "memory", "schema_records"}
required = {"id", "record_key", "kind", "content_hash", "metadata", "embedding", "indexed_at"}
for table in tables:
assert {row.column_name for row in rows if row.table_name == table} == required
assert vector_types == [
("evidence", "vector(768)"),
("memory", "vector(768)"),
("schema_records", "vector(768)"),
]
def test_roles_have_runtime_privileges_only(database_url):
from tht.cli.vector_migrate_cmd import migrate
migrate(database_url)
admin = create_engine(database_url)
with admin.begin() as connection:
connection.exec_driver_sql("ALTER ROLE vector_reader LOGIN PASSWORD 'reader-test-only'")
connection.exec_driver_sql("ALTER ROLE vector_writer LOGIN PASSWORD 'writer-test-only'")
url = admin.url
reader = create_engine(url.set(username="vector_reader", password="reader-test-only"))
writer = create_engine(url.set(username="vector_writer", password="writer-test-only"))
from tht.adapters.vector.pgvector import PgVectorStore
common = {
"host": url.host,
"port": url.port,
"database": url.database,
"schema": "vectors",
}
reader_config = DatabaseConfig(
**common, user="vector_reader", password="reader-test-only"
)
writer_config = DatabaseConfig(
**common, user="vector_writer", password="writer-test-only"
)
store = PgVectorStore(reader_config, writer_config, expected_dimension=768)
assert store.health().ok is True
assert store.upsert(
"memory",
[
VectorWriteRecord(
record=VectorRecord(
id="adapter-write",
kind="memory",
ref="session:test",
title="test",
content="test",
),
embedding=[0.0] * 768,
content_hash="adapter-hash",
)
],
) == 1
with reader.connect() as connection:
connection.execute(text("SELECT metadata, embedding FROM vectors.memory")).all()
with pytest.raises(ProgrammingError):
with reader.begin() as connection:
connection.execute(
text(
"INSERT INTO vectors.memory "
"(record_key, kind, content_hash, metadata, embedding) "
"VALUES ('reader-write', 'memory', 'x', '{}', "
"array_fill(0, ARRAY[768])::vector)"
)
)
with writer.begin() as connection:
connection.execute(
text(
"INSERT INTO vectors.memory "
"(record_key, kind, content_hash, metadata, embedding) "
"VALUES ('writer-ok', 'memory', 'x', '{}', array_fill(0, ARRAY[768])::vector)"
)
)
assert connection.execute(
text("SELECT content_hash FROM vectors.memory WHERE record_key = 'writer-ok'")
).scalar_one() == "x"
connection.execute(
text("UPDATE vectors.memory SET content_hash = 'y' WHERE record_key = 'writer-ok'")
)
with pytest.raises(ProgrammingError):
with writer.connect() as connection:
connection.execute(text("SELECT metadata FROM vectors.memory")).all()
with pytest.raises(ProgrammingError):
with writer.begin() as connection:
connection.execute(text("DELETE FROM vectors.memory WHERE record_key = 'writer-ok'"))
reader.dispose()
writer.dispose()
admin.dispose()
def test_status_json_is_pristine(database_url, monkeypatch):
monkeypatch.setenv("THT_VECTOR_ADMIN_URL", database_url)
result = CliRunner().invoke(app, ["vector", "migrate", "--status", "--json"])
assert result.exit_code == 0, result.output
assert json.loads(result.stdout) == {
"applied": ["001", "002", "003"],
"drifted": [],
"pending": [],
}
assert result.stderr == ""
def test_checksum_drift_is_reported_and_refused(database_url, tmp_path):
from tht.cli.vector_migrate_cmd import MigrationError, migrate, migration_status
migrations = _copy_migrations(tmp_path)
migrate(database_url, migrations)
(migrations / "002_schema_tables.sql").write_text("SELECT 2;\n")
assert [item.version for item in migration_status(database_url, migrations).drifted] == [
"002"
]
with pytest.raises(MigrationError, match="checksum drift"):
migrate(database_url, migrations)
def test_failed_batch_rolls_back_schema_and_ledger(database_url, tmp_path):
from tht.cli.vector_migrate_cmd import MigrationError, migrate, migration_status
migrations = tmp_path / "failed"
migrations.mkdir()
(migrations / "101_first.sql").write_text("CREATE TABLE public.must_rollback (id int);\n")
(migrations / "102_broken.sql").write_text("THIS IS NOT SQL;\n")
with pytest.raises(MigrationError, match="102_broken.sql"):
migrate(database_url, migrations)
engine = create_engine(database_url)
with engine.connect() as connection:
assert connection.execute(text("SELECT to_regclass('public.must_rollback')")).scalar() is None
engine.dispose()
status = migration_status(database_url, migrations)
assert status.applied == ()
assert [item.version for item in status.pending] == ["101", "102"]
def _copy_migrations(tmp_path: Path) -> Path:
source = Path(__file__).parents[2] / "migrations" / "vector"
target = tmp_path / "migrations"
target.mkdir()
for migration in source.glob("*.sql"):
(target / migration.name).write_bytes(migration.read_bytes())
return target