"""Postgres migration runner — applies db/pg_migrations/*.sql in order. Mirrors db/migrate.py: ordered .sql files tracked in a `_pg_migrations` table so re-running is idempotent. Uses an asyncpg pool. Retries on connection failure (3 attempts, 2s backoff — R-MT-02 mitigation). """ from __future__ import annotations import asyncio import datetime as _dt from pathlib import Path import asyncpg _DEFAULT_MIGRATIONS_DIR = Path(__file__).resolve().parent / "pg_migrations" _RETRY_ATTEMPTS = 3 _RETRY_BACKOFF_S = 2.0 async def apply_pg_migrations( pool: asyncpg.Pool, migrations_dir: Path | None = None, ) -> list[str]: """Apply all pending Postgres migrations in order. Returns applied names. Idempotent — no-op if all migrations are already applied. Each migration runs within a transaction; the `_pg_migrations` tracking row is inserted in the same transaction so a failure rolls back cleanly. """ mdir = migrations_dir or _DEFAULT_MIGRATIONS_DIR if not mdir.exists(): return [] async def _run() -> list[str]: async with pool.acquire() as conn: await conn.execute( "CREATE TABLE IF NOT EXISTS _pg_migrations (" "id TEXT PRIMARY KEY, applied_at TIMESTAMPTZ NOT NULL DEFAULT now()" ")" ) rows = await conn.fetch("SELECT id FROM _pg_migrations") applied_ids = {r["id"] for r in rows} applied: list[str] = [] for sql_path in sorted(mdir.glob("*.sql")): mid = sql_path.stem if mid in applied_ids: continue sql = sql_path.read_text(encoding="utf-8") async with conn.transaction(): await conn.execute(sql) await conn.execute( "INSERT INTO _pg_migrations (id) VALUES ($1)", mid ) applied.append(mid) return applied last_exc: Exception | None = None for attempt in range(1, _RETRY_ATTEMPTS + 1): try: return await _run() except (asyncpg.PostgresConnectionError, ConnectionError, OSError) as exc: last_exc = exc if attempt < _RETRY_ATTEMPTS: await asyncio.sleep(_RETRY_BACKOFF_S) continue assert last_exc is not None raise last_exc __all__ = ["apply_pg_migrations"]