304 lines
12 KiB
Python
304 lines
12 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", "vectors.vector(768)"),
|
|
("memory", "vectors.vector(768)"),
|
|
("schema_records", "vectors.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])::vectors.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])::vectors.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_unknown_applied_version_is_downgrade_drift(database_url):
|
|
from tht.cli.vector_migrate_cmd import MigrationError, migrate, migration_status
|
|
|
|
migrate(database_url)
|
|
engine = create_engine(database_url)
|
|
with engine.begin() as connection:
|
|
connection.execute(
|
|
text(
|
|
"INSERT INTO public.tht_vector_migrations (version, name, checksum) "
|
|
"VALUES ('999', 'future', 'future-checksum'), "
|
|
"('future_x', 'future_named', 'future-checksum')"
|
|
)
|
|
)
|
|
try:
|
|
with pytest.raises(
|
|
MigrationError, match="absent from local manifest: 999, future_x"
|
|
):
|
|
migration_status(database_url)
|
|
with pytest.raises(
|
|
MigrationError, match="absent from local manifest: 999, future_x"
|
|
):
|
|
migrate(database_url)
|
|
finally:
|
|
with engine.begin() as connection:
|
|
connection.execute(
|
|
text(
|
|
"DELETE FROM public.tht_vector_migrations "
|
|
"WHERE version IN ('999', 'future_x')"
|
|
)
|
|
)
|
|
engine.dispose()
|
|
|
|
|
|
def test_migration_versions_sort_numerically_and_reject_numeric_duplicates(tmp_path):
|
|
from tht.cli.vector_migrate_cmd import MigrationError, _discover
|
|
|
|
migrations = tmp_path / "ordered"
|
|
migrations.mkdir()
|
|
(migrations / "10_tenth.sql").write_text("SELECT 10;\n")
|
|
(migrations / "2_second.sql").write_text("SELECT 2;\n")
|
|
assert [item.version for item in _discover(migrations)] == ["2", "10"]
|
|
|
|
(migrations / "02_duplicate.sql").write_text("SELECT 2;\n")
|
|
with pytest.raises(MigrationError, match="Duplicate migration version: 2"):
|
|
_discover(migrations)
|
|
|
|
|
|
def test_hostile_admin_search_path_cannot_shadow_migration_objects(database_url):
|
|
from tht.cli.vector_migrate_cmd import migrate
|
|
|
|
admin = create_engine(database_url, isolation_level="AUTOCOMMIT")
|
|
with admin.connect() as connection:
|
|
connection.exec_driver_sql("DROP DATABASE IF EXISTS vector_hostile")
|
|
connection.exec_driver_sql("CREATE DATABASE vector_hostile")
|
|
hostile_url = admin.url.set(database="vector_hostile")
|
|
hostile = create_engine(hostile_url)
|
|
try:
|
|
with hostile.begin() as connection:
|
|
connection.exec_driver_sql("CREATE SCHEMA shadow")
|
|
connection.exec_driver_sql(
|
|
"CREATE TABLE shadow.tht_vector_migrations "
|
|
"(version text, checksum text, poisoned boolean DEFAULT true)"
|
|
)
|
|
connection.exec_driver_sql("ALTER ROLE test SET search_path = shadow, public")
|
|
hostile.dispose()
|
|
|
|
migrate(hostile_url.render_as_string(hide_password=False))
|
|
|
|
verification = create_engine(hostile_url)
|
|
with verification.connect() as connection:
|
|
assert connection.execute(
|
|
text("SELECT count(*) FROM public.tht_vector_migrations")
|
|
).scalar_one() == 3
|
|
assert connection.execute(
|
|
text("SELECT count(*) FROM shadow.tht_vector_migrations")
|
|
).scalar_one() == 0
|
|
assert connection.execute(
|
|
text(
|
|
"SELECT format_type(a.atttypid, a.atttypmod) "
|
|
"FROM pg_catalog.pg_attribute a "
|
|
"WHERE a.attrelid = 'vectors.memory'::pg_catalog.regclass "
|
|
"AND a.attname = 'embedding'"
|
|
)
|
|
).scalar_one() == "vectors.vector(768)"
|
|
verification.dispose()
|
|
finally:
|
|
cleanup = create_engine(database_url, isolation_level="AUTOCOMMIT")
|
|
with cleanup.connect() as connection:
|
|
connection.exec_driver_sql("ALTER ROLE test RESET search_path")
|
|
connection.exec_driver_sql(
|
|
"SELECT pg_catalog.pg_terminate_backend(pid) FROM pg_catalog.pg_stat_activity "
|
|
"WHERE datname = 'vector_hostile' AND pid <> pg_catalog.pg_backend_pid()"
|
|
)
|
|
connection.exec_driver_sql("DROP DATABASE IF EXISTS vector_hostile")
|
|
cleanup.dispose()
|
|
admin.dispose()
|
|
|
|
|
|
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 = _copy_migrations(tmp_path)
|
|
(migrations / "004_first.sql").write_text("CREATE TABLE public.must_rollback (id int);\n")
|
|
(migrations / "005_broken.sql").write_text("THIS IS NOT SQL;\n")
|
|
|
|
with pytest.raises(MigrationError, match="005_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 [item.version for item in status.applied] == ["001", "002", "003"]
|
|
assert [item.version for item in status.pending] == ["004", "005"]
|
|
|
|
|
|
def _copy_migrations(tmp_path: Path) -> Path:
|
|
source = Path(__file__).parents[2] / "tht" / "migrations" / "vector"
|
|
target = tmp_path / "migrations"
|
|
target.mkdir()
|
|
for migration in source.glob("*.sql"):
|
|
(target / migration.name).write_bytes(migration.read_bytes())
|
|
return target
|