* fix(migrations): use HINDSIGHT_API_MIGRATION_DATABASE_URL when set Session-level advisory locks are broken when the database URL goes through PgBouncer in transaction mode: the backend connection is returned to the pool on COMMIT, orphaning the lock, so multiple pods can simultaneously run migrations for the same schema. When HINDSIGHT_API_MIGRATION_DATABASE_URL is set, use it for both the advisory lock connection and the Alembic run. Callers should point this at the direct PostgreSQL endpoint (bypassing the pooler) so the session-level lock is held for the full migration duration. * refactor(migrations): move MIGRATION_DATABASE_URL to standard config Wire HINDSIGHT_API_MIGRATION_DATABASE_URL through HindsightConfig instead of reading os.getenv() directly in migrations.py. Add the field to the dataclass, from_env(), log_config(), all call sites, .env.example, and the configuration docs page. * fix: update test mocks for migration_database_url kwarg and regenerate docs skill --------- Co-authored-by: Nicolò Boschi <boschi1997@gmail.com>
487 lines
19 KiB
Python
487 lines
19 KiB
Python
"""
|
|
Tests for admin backup and restore functionality.
|
|
|
|
These tests use an isolated schema to avoid interfering with other tests.
|
|
The backup/restore operations truncate tables, which would cause deadlocks
|
|
and race conditions if run against the shared public schema.
|
|
"""
|
|
|
|
import tempfile
|
|
import uuid
|
|
import zipfile
|
|
from pathlib import Path
|
|
|
|
import asyncpg
|
|
import pytest
|
|
import pytest_asyncio
|
|
|
|
import hindsight_api.admin.cli as admin_cli
|
|
from hindsight_api.admin.cli import _backup, _restore, BACKUP_TABLES
|
|
from hindsight_api.extensions import Tenant
|
|
from hindsight_api.migrations import run_migrations
|
|
|
|
|
|
# Run these tests sequentially since they do full DB backup/restore
|
|
pytestmark = pytest.mark.xdist_group(name="backup_restore")
|
|
|
|
|
|
@pytest_asyncio.fixture(scope="function")
|
|
async def backup_test_schema(pg0_db_url, embeddings):
|
|
"""Create an isolated schema for backup/restore tests.
|
|
|
|
Uses a unique schema name per test invocation to avoid conflicts with
|
|
parallel test runs or leftover state from interrupted runs.
|
|
|
|
Returns a tuple of (db_url, schema_name, fq_helper, embeddings).
|
|
"""
|
|
# Initialize embeddings if not already done
|
|
await embeddings.initialize()
|
|
|
|
# Use unique schema name to avoid conflicts
|
|
schema_name = f"backup_test_{uuid.uuid4().hex[:8]}"
|
|
|
|
def _fq(table: str) -> str:
|
|
"""Get fully-qualified table name in test schema."""
|
|
return f"{schema_name}.{table}"
|
|
|
|
conn = await asyncpg.connect(pg0_db_url)
|
|
try:
|
|
await conn.execute(f"CREATE SCHEMA {schema_name}")
|
|
finally:
|
|
await conn.close()
|
|
|
|
# Run migrations on the isolated schema
|
|
run_migrations(pg0_db_url, schema=schema_name)
|
|
|
|
yield pg0_db_url, schema_name, _fq, embeddings
|
|
|
|
# Cleanup after test
|
|
conn = await asyncpg.connect(pg0_db_url)
|
|
try:
|
|
await conn.execute(f"DROP SCHEMA IF EXISTS {schema_name} CASCADE")
|
|
finally:
|
|
await conn.close()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_backup_restore_roundtrip(backup_test_schema):
|
|
"""Test that backup and restore preserves all data correctly."""
|
|
db_url, schema_name, _fq, embeddings = backup_test_schema
|
|
bank_id = f"test-backup-{uuid.uuid4().hex[:8]}"
|
|
conn = await asyncpg.connect(db_url)
|
|
|
|
try:
|
|
# Create a bank
|
|
await conn.execute(
|
|
f"INSERT INTO {_fq('banks')} (bank_id) VALUES ($1) ON CONFLICT DO NOTHING",
|
|
bank_id,
|
|
)
|
|
|
|
# Create some test memory units with embeddings
|
|
# Convert embedding list to pgvector format string
|
|
embedding_list = embeddings.encode(["Test content about Alice"])[0]
|
|
embedding_str = "[" + ",".join(str(x) for x in embedding_list) + "]"
|
|
for text in [
|
|
"Alice is a software engineer who loves Python.",
|
|
"Bob works with Alice on the backend team.",
|
|
"The team uses PostgreSQL for their database.",
|
|
]:
|
|
await conn.execute(
|
|
f"""INSERT INTO {_fq('memory_units')}
|
|
(bank_id, text, fact_type, embedding, event_date)
|
|
VALUES ($1, $2, 'world', $3::vector, NOW())""",
|
|
bank_id,
|
|
text,
|
|
embedding_str,
|
|
)
|
|
|
|
# Get counts before backup
|
|
counts_before = {}
|
|
for table in BACKUP_TABLES:
|
|
counts_before[table] = await conn.fetchval(f"SELECT COUNT(*) FROM {_fq(table)}")
|
|
|
|
# Verify we have data
|
|
assert counts_before["banks"] > 0
|
|
assert counts_before["memory_units"] > 0
|
|
|
|
finally:
|
|
await conn.close()
|
|
|
|
# Backup to a temp file
|
|
with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as f:
|
|
backup_path = Path(f.name)
|
|
|
|
try:
|
|
manifest = await _backup(db_url, backup_path, schema=schema_name)
|
|
|
|
# Verify backup file exists and is valid
|
|
assert backup_path.exists()
|
|
assert backup_path.stat().st_size > 0
|
|
|
|
# Verify manifest
|
|
assert manifest["version"] == "1"
|
|
assert "created_at" in manifest
|
|
for table in BACKUP_TABLES:
|
|
assert table in manifest["tables"]
|
|
assert manifest["tables"][table]["rows"] == counts_before[table]
|
|
|
|
# Verify zip contents
|
|
with zipfile.ZipFile(backup_path, "r") as zf:
|
|
assert "manifest.json" in zf.namelist()
|
|
for table in BACKUP_TABLES:
|
|
assert f"{table}.bin" in zf.namelist()
|
|
|
|
# Clear all data
|
|
conn = await asyncpg.connect(db_url)
|
|
try:
|
|
for table in reversed(BACKUP_TABLES):
|
|
await conn.execute(f"TRUNCATE TABLE {_fq(table)} CASCADE")
|
|
|
|
# Verify data is gone
|
|
for table in BACKUP_TABLES:
|
|
count = await conn.fetchval(f"SELECT COUNT(*) FROM {_fq(table)}")
|
|
assert count == 0, f"Table {table} should be empty after truncate"
|
|
finally:
|
|
await conn.close()
|
|
|
|
# Restore from backup
|
|
await _restore(db_url, backup_path, schema=schema_name)
|
|
|
|
# Verify counts match original
|
|
conn = await asyncpg.connect(db_url)
|
|
try:
|
|
for table in BACKUP_TABLES:
|
|
count = await conn.fetchval(f"SELECT COUNT(*) FROM {_fq(table)}")
|
|
assert count == counts_before[table], f"Table {table} count mismatch after restore"
|
|
|
|
# Verify data content is preserved
|
|
texts = await conn.fetch(
|
|
f"SELECT text FROM {_fq('memory_units')} WHERE bank_id = $1",
|
|
bank_id,
|
|
)
|
|
text_content = " ".join(r["text"] for r in texts)
|
|
assert "Alice" in text_content or "software" in text_content
|
|
finally:
|
|
await conn.close()
|
|
|
|
finally:
|
|
# Cleanup
|
|
if backup_path.exists():
|
|
backup_path.unlink()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_backup_restore_preserves_all_column_types(backup_test_schema):
|
|
"""Test that all column types are preserved: vectors, UUIDs, timestamps, JSONB."""
|
|
db_url, schema_name, _fq, embeddings = backup_test_schema
|
|
bank_id = f"test-types-{uuid.uuid4().hex[:8]}"
|
|
conn = await asyncpg.connect(db_url)
|
|
|
|
try:
|
|
# Create a bank
|
|
await conn.execute(
|
|
f"INSERT INTO {_fq('banks')} (bank_id) VALUES ($1) ON CONFLICT DO NOTHING",
|
|
bank_id,
|
|
)
|
|
|
|
# Create a memory unit with all column types
|
|
# Convert embedding list to pgvector format string
|
|
embedding_list = embeddings.encode(["John Smith engineer"])[0]
|
|
embedding_str = "[" + ",".join(str(x) for x in embedding_list) + "]"
|
|
await conn.execute(
|
|
f"""INSERT INTO {_fq('memory_units')}
|
|
(bank_id, text, fact_type, embedding, event_date, metadata)
|
|
VALUES ($1, $2, 'world', $3::vector, NOW(), $4)""",
|
|
bank_id,
|
|
"John Smith is a senior engineer at Acme Corp since 2020.",
|
|
embedding_str,
|
|
'{"key": "value"}',
|
|
)
|
|
|
|
# Create an entity
|
|
await conn.execute(
|
|
f"""INSERT INTO {_fq('entities')}
|
|
(bank_id, canonical_name, metadata)
|
|
VALUES ($1, $2, $3)""",
|
|
bank_id,
|
|
"John Smith",
|
|
'{"role": "engineer"}',
|
|
)
|
|
|
|
# Get original data
|
|
original_unit = await conn.fetchrow(
|
|
f"""SELECT id, embedding, event_date, created_at, metadata, text
|
|
FROM {_fq('memory_units')} WHERE bank_id = $1 LIMIT 1""",
|
|
bank_id,
|
|
)
|
|
original_entity = await conn.fetchrow(
|
|
f"""SELECT id, first_seen, last_seen, metadata, canonical_name
|
|
FROM {_fq('entities')} WHERE bank_id = $1 LIMIT 1""",
|
|
bank_id,
|
|
)
|
|
original_bank = await conn.fetchrow(
|
|
f"SELECT bank_id, created_at, updated_at FROM {_fq('banks')} WHERE bank_id = $1",
|
|
bank_id,
|
|
)
|
|
finally:
|
|
await conn.close()
|
|
|
|
assert original_unit is not None, "Should have created memory units"
|
|
assert original_unit["embedding"] is not None, "Should have embedding"
|
|
assert original_unit["id"] is not None, "Should have UUID"
|
|
assert original_entity is not None, "Should have created entities"
|
|
|
|
with tempfile.NamedTemporaryFile(suffix=".zip", delete=False) as f:
|
|
backup_path = Path(f.name)
|
|
|
|
try:
|
|
await _backup(db_url, backup_path, schema=schema_name)
|
|
|
|
# Clear all data
|
|
conn = await asyncpg.connect(db_url)
|
|
try:
|
|
for table in reversed(BACKUP_TABLES):
|
|
await conn.execute(f"TRUNCATE TABLE {_fq(table)} CASCADE")
|
|
finally:
|
|
await conn.close()
|
|
|
|
await _restore(db_url, backup_path, schema=schema_name)
|
|
|
|
# Verify all column types are preserved exactly
|
|
conn = await asyncpg.connect(db_url)
|
|
try:
|
|
restored_unit = await conn.fetchrow(
|
|
f"""SELECT id, embedding, event_date, created_at, metadata, text
|
|
FROM {_fq('memory_units')} WHERE bank_id = $1 LIMIT 1""",
|
|
bank_id,
|
|
)
|
|
restored_entity = await conn.fetchrow(
|
|
f"""SELECT id, first_seen, last_seen, metadata, canonical_name
|
|
FROM {_fq('entities')} WHERE bank_id = $1 LIMIT 1""",
|
|
bank_id,
|
|
)
|
|
restored_bank = await conn.fetchrow(
|
|
f"SELECT bank_id, created_at, updated_at FROM {_fq('banks')} WHERE bank_id = $1",
|
|
bank_id,
|
|
)
|
|
finally:
|
|
await conn.close()
|
|
|
|
# Verify memory_units
|
|
assert restored_unit is not None, "Should have restored memory unit"
|
|
assert restored_unit["id"] == original_unit["id"], "UUID should match exactly"
|
|
assert restored_unit["text"] == original_unit["text"], "Text should match"
|
|
assert list(restored_unit["embedding"]) == list(original_unit["embedding"]), "Vector embedding should match exactly"
|
|
assert restored_unit["event_date"] == original_unit["event_date"], "Timestamp should match exactly"
|
|
assert restored_unit["created_at"] == original_unit["created_at"], "Created timestamp should match"
|
|
assert restored_unit["metadata"] == original_unit["metadata"], "JSONB metadata should match"
|
|
|
|
# Verify entities
|
|
assert restored_entity is not None, "Should have restored entity"
|
|
assert restored_entity["id"] == original_entity["id"], "Entity UUID should match"
|
|
assert restored_entity["canonical_name"] == original_entity["canonical_name"], "Entity name should match"
|
|
assert restored_entity["first_seen"] == original_entity["first_seen"], "Entity first_seen should match"
|
|
assert restored_entity["last_seen"] == original_entity["last_seen"], "Entity last_seen should match"
|
|
assert restored_entity["metadata"] == original_entity["metadata"], "Entity metadata should match"
|
|
|
|
# Verify banks
|
|
assert restored_bank is not None, "Should have restored bank"
|
|
assert restored_bank["bank_id"] == original_bank["bank_id"], "Bank ID should match"
|
|
assert restored_bank["created_at"] == original_bank["created_at"], "Bank created_at should match"
|
|
|
|
finally:
|
|
if backup_path.exists():
|
|
backup_path.unlink()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_migration_without_schema_discovers_and_deduplicates_schemas(monkeypatch):
|
|
"""run-db-migration without --schema should include the base schema and deduplicate tenant schemas."""
|
|
calls: dict[str, list] = {
|
|
"run_migrations": [],
|
|
"ensure_vector_extension": [],
|
|
"ensure_text_search_extension": [],
|
|
}
|
|
|
|
class MockTenantExtension:
|
|
async def list_tenants(self):
|
|
return [
|
|
Tenant(schema="public"),
|
|
Tenant(schema="tenant_demo"),
|
|
Tenant(schema="tenant_demo"),
|
|
]
|
|
|
|
async def fake_resolve_database_url(db_url: str) -> str:
|
|
return f"resolved::{db_url}"
|
|
|
|
def fake_run_migrations(database_url: str, schema: str | None = None, **kwargs) -> None:
|
|
calls["run_migrations"].append((database_url, schema))
|
|
|
|
def fake_ensure_vector_extension(
|
|
database_url: str,
|
|
vector_extension: str = "pgvector",
|
|
schema: str | None = None,
|
|
) -> None:
|
|
calls["ensure_vector_extension"].append((database_url, vector_extension, schema))
|
|
|
|
def fake_ensure_text_search_extension(
|
|
database_url: str,
|
|
text_search_extension: str = "native",
|
|
schema: str | None = None,
|
|
) -> None:
|
|
calls["ensure_text_search_extension"].append((database_url, text_search_extension, schema))
|
|
|
|
monkeypatch.setenv("HINDSIGHT_API_DATABASE_URL", "postgresql://test")
|
|
monkeypatch.setattr(admin_cli, "load_extension", lambda *args, **kwargs: MockTenantExtension())
|
|
monkeypatch.setattr(admin_cli, "resolve_database_url", fake_resolve_database_url)
|
|
|
|
from hindsight_api import migrations as migrations_module
|
|
|
|
monkeypatch.setattr(migrations_module, "run_migrations", fake_run_migrations)
|
|
monkeypatch.setattr(migrations_module, "ensure_vector_extension", fake_ensure_vector_extension)
|
|
monkeypatch.setattr(migrations_module, "ensure_text_search_extension", fake_ensure_text_search_extension)
|
|
|
|
schemas = await admin_cli._run_migration("postgresql://test")
|
|
|
|
assert schemas == ["public", "tenant_demo"]
|
|
assert calls["run_migrations"] == [
|
|
("resolved::postgresql://test", "public"),
|
|
("resolved::postgresql://test", "tenant_demo"),
|
|
]
|
|
assert calls["ensure_vector_extension"] == [
|
|
("resolved::postgresql://test", "pgvector", "public"),
|
|
("resolved::postgresql://test", "pgvector", "tenant_demo"),
|
|
]
|
|
assert calls["ensure_text_search_extension"] == [
|
|
("resolved::postgresql://test", "native", "public"),
|
|
("resolved::postgresql://test", "native", "tenant_demo"),
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_migration_without_schema_runs_optional_post_migration_hooks(monkeypatch):
|
|
"""Embedding dimension sync should be optional, while vector/text checks always run."""
|
|
monkeypatch.setenv("HINDSIGHT_API_DATABASE_URL", "postgresql://test")
|
|
calls: dict[str, list] = {
|
|
"run_migrations": [],
|
|
"ensure_embedding_dimension": [],
|
|
"ensure_vector_extension": [],
|
|
"ensure_text_search_extension": [],
|
|
}
|
|
|
|
class MockTenantExtension:
|
|
async def list_tenants(self):
|
|
return [Tenant(schema="tenant_demo")]
|
|
|
|
async def fake_resolve_database_url(db_url: str) -> str:
|
|
return f"resolved::{db_url}"
|
|
|
|
def fake_run_migrations(database_url: str, schema: str | None = None, **kwargs) -> None:
|
|
calls["run_migrations"].append((database_url, schema))
|
|
|
|
def fake_ensure_embedding_dimension(
|
|
database_url: str,
|
|
dimension: int,
|
|
schema: str | None = None,
|
|
vector_extension: str = "pgvector",
|
|
) -> None:
|
|
calls["ensure_embedding_dimension"].append((database_url, dimension, schema, vector_extension))
|
|
|
|
def fake_ensure_vector_extension(
|
|
database_url: str,
|
|
vector_extension: str = "pgvector",
|
|
schema: str | None = None,
|
|
) -> None:
|
|
calls["ensure_vector_extension"].append((database_url, vector_extension, schema))
|
|
|
|
def fake_ensure_text_search_extension(
|
|
database_url: str,
|
|
text_search_extension: str = "native",
|
|
schema: str | None = None,
|
|
) -> None:
|
|
calls["ensure_text_search_extension"].append((database_url, text_search_extension, schema))
|
|
|
|
monkeypatch.setattr(admin_cli, "load_extension", lambda *args, **kwargs: MockTenantExtension())
|
|
monkeypatch.setattr(admin_cli, "resolve_database_url", fake_resolve_database_url)
|
|
|
|
from hindsight_api import migrations as migrations_module
|
|
|
|
monkeypatch.setattr(migrations_module, "run_migrations", fake_run_migrations)
|
|
monkeypatch.setattr(migrations_module, "ensure_embedding_dimension", fake_ensure_embedding_dimension)
|
|
monkeypatch.setattr(migrations_module, "ensure_vector_extension", fake_ensure_vector_extension)
|
|
monkeypatch.setattr(migrations_module, "ensure_text_search_extension", fake_ensure_text_search_extension)
|
|
|
|
schemas = await admin_cli._run_migration(
|
|
"postgresql://test",
|
|
base_schema="public",
|
|
embedding_dimension=384,
|
|
)
|
|
|
|
assert schemas == ["public", "tenant_demo"]
|
|
assert calls["run_migrations"] == [
|
|
("resolved::postgresql://test", "public"),
|
|
("resolved::postgresql://test", "tenant_demo"),
|
|
]
|
|
assert calls["ensure_embedding_dimension"] == [
|
|
("resolved::postgresql://test", 384, "public", "pgvector"),
|
|
("resolved::postgresql://test", 384, "tenant_demo", "pgvector"),
|
|
]
|
|
assert calls["ensure_vector_extension"] == [
|
|
("resolved::postgresql://test", "pgvector", "public"),
|
|
("resolved::postgresql://test", "pgvector", "tenant_demo"),
|
|
]
|
|
assert calls["ensure_text_search_extension"] == [
|
|
("resolved::postgresql://test", "native", "public"),
|
|
("resolved::postgresql://test", "native", "tenant_demo"),
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_run_migration_with_schema_only_runs_requested_schema(monkeypatch):
|
|
"""run-db-migration with --schema should only migrate the requested schema."""
|
|
monkeypatch.setenv("HINDSIGHT_API_DATABASE_URL", "postgresql://test")
|
|
calls: dict[str, list] = {
|
|
"run_migrations": [],
|
|
"ensure_vector_extension": [],
|
|
"ensure_text_search_extension": [],
|
|
}
|
|
|
|
class MockTenantExtension:
|
|
async def list_tenants(self):
|
|
return [Tenant(schema="tenant_demo"), Tenant(schema="tenant_other")]
|
|
|
|
async def fake_resolve_database_url(db_url: str) -> str:
|
|
return f"resolved::{db_url}"
|
|
|
|
def fake_run_migrations(database_url: str, schema: str | None = None, **kwargs) -> None:
|
|
calls["run_migrations"].append((database_url, schema))
|
|
|
|
def fake_ensure_vector_extension(
|
|
database_url: str,
|
|
vector_extension: str = "pgvector",
|
|
schema: str | None = None,
|
|
) -> None:
|
|
calls["ensure_vector_extension"].append((database_url, vector_extension, schema))
|
|
|
|
def fake_ensure_text_search_extension(
|
|
database_url: str,
|
|
text_search_extension: str = "native",
|
|
schema: str | None = None,
|
|
) -> None:
|
|
calls["ensure_text_search_extension"].append((database_url, text_search_extension, schema))
|
|
|
|
monkeypatch.setattr(admin_cli, "load_extension", lambda *args, **kwargs: MockTenantExtension())
|
|
monkeypatch.setattr(admin_cli, "resolve_database_url", fake_resolve_database_url)
|
|
|
|
from hindsight_api import migrations as migrations_module
|
|
|
|
monkeypatch.setattr(migrations_module, "run_migrations", fake_run_migrations)
|
|
monkeypatch.setattr(migrations_module, "ensure_vector_extension", fake_ensure_vector_extension)
|
|
monkeypatch.setattr(migrations_module, "ensure_text_search_extension", fake_ensure_text_search_extension)
|
|
|
|
schemas = await admin_cli._run_migration("postgresql://test", schema="tenant_demo")
|
|
|
|
assert schemas == ["tenant_demo"]
|
|
assert calls["run_migrations"] == [("resolved::postgresql://test", "tenant_demo")]
|
|
assert calls["ensure_vector_extension"] == [("resolved::postgresql://test", "pgvector", "tenant_demo")]
|
|
assert calls["ensure_text_search_extension"] == [("resolved::postgresql://test", "native", "tenant_demo")]
|