diff --git a/tools/migrate_v45_to_v46.py b/tools/migrate_v45_to_v46.py new file mode 100644 index 0000000..90bc781 --- /dev/null +++ b/tools/migrate_v45_to_v46.py @@ -0,0 +1,329 @@ +"""One-time migration of a V4.5 database into the re-leveled V4.6 schema. + +V4.6 re-levels the schema from the current SQLModel metadata rather than +running a chain of hand-rolled upgrade functions. The column sets are +unchanged; what changed is index coverage ([HIGH-04]), the ``use_alter`` +break in the ``source``/``execution_attempt`` foreign key cycle, and the +relationship loading strategy ([CRIT-02]). This script therefore performs a +faithful, foreign-key-ordered row copy. + +Design notes: + +- The backup is read with plain ``sqlite3`` rather than through the ORM. The + V4.5 file is not guaranteed to satisfy the V4.6 mappers, and reading raw + rows means no relationship is ever traversed, so ``lazy="raise"`` cannot + bite. +- The target is written through SQLAlchemy Core against the live metadata, so + the same script works against PostgreSQL when that cutover happens. +- Identity is preserved exactly: UUIDs, digests, timestamps, attempt numbers, + and ``preferred_execution_attempt_id`` selections carry across unchanged. + No evidence payload is reinterpreted, normalized, or regenerated. +- No on-disk Source file, portrait, or artifact file is read for writing or + modified. ``--verify-artifacts`` reads artifact files, but only to hash + them. +- The script is idempotent: a row whose primary key already exists in the + target is skipped, never updated. It is never invoked from application + startup and never runs in the test suite. + +Usage:: + + python tools/migrate_v45_to_v46.py --dry-run + python tools/migrate_v45_to_v46.py --verify-artifacts +""" + +from __future__ import annotations + +import argparse +import json +import sqlite3 +import sys +from collections.abc import Iterator +from collections.abc import Sequence +from datetime import date +from datetime import datetime +from pathlib import Path +from typing import Any +from uuid import UUID + +from sqlalchemy import Column +from sqlalchemy import Table +from sqlalchemy import create_engine +from sqlalchemy import insert +from sqlalchemy import inspect as sqlalchemy_inspect +from sqlalchemy import select +from sqlalchemy import update +from sqlalchemy.engine import Connection +from sqlmodel import SQLModel + +from transcription.config import Settings +from transcription.config import get_settings +from transcription.db import models as _models # noqa: F401 (registers every table) +from transcription.db.engine import get_database_url + +DEFAULT_BACKUP = Path("data/transcription.db.pre-v46.bak") + +#: ``source.preferred_execution_attempt_id`` points at ``execution_attempt``, +#: which points back at ``source``. The cycle is broken with ``use_alter`` in +#: the metadata, so ``source`` rows are inserted with the column cleared and +#: the selections are replayed once ``execution_attempt`` is populated. +DEFERRED_TABLE = "source" +DEFERRED_COLUMN = "preferred_execution_attempt_id" + +#: Row counts the V4.5 backup is expected to carry, used as a pre-flight guard +#: so the script cannot silently run against the wrong file. +EXPECTED_SOURCE_COUNTS = { + "document": 8, + "document_person": 11, + "document_type": 7, + "execution_attempt": 80, + "job": 11, + "job_source": 79, + "person": 5, + "person_role": 3, + "processing_artifact": 2, + "source": 76, +} + + +def _coerce(column: Column[Any], value: object) -> object: + """Convert a raw SQLite value into what the target column's type binds. + + SQLite hands back strings and integers; the V4.6 columns bind ``UUID``, + ``datetime``, ``date``, ``bool``, enum members, and decoded JSON. The + conversion is lossless in both directions. + """ + if value is None: + return None + + match type(column.type).__name__: + case "Uuid": + return value if isinstance(value, UUID) else UUID(str(value)) + case "DateTime": + return value if isinstance(value, datetime) else datetime.fromisoformat(str(value)) + case "Date": + return value if isinstance(value, date) else date.fromisoformat(str(value)) + case "Boolean": + return bool(value) + case "JSONBCompat": + if isinstance(value, str | bytes | bytearray): + return json.loads(value) + return value + case "Enum": + enum_class = getattr(column.type, "enum_class", None) + if enum_class is None: + return value + # The same JobSourceStatus enum is persisted by value on + # job_source.status and by name on execution_attempt.status, + # because only the former declares values_callable. Accept either + # spelling so the copy round-trips both columns faithfully. + try: + return enum_class(value) + except ValueError: + return enum_class[str(value)] + case _: + return value + + +def _read_table(backup: sqlite3.Connection, table: Table) -> list[dict[str, object]]: + """Read every row of ``table`` from the backup, coerced for the target.""" + names = [column.name for column in table.columns] + quoted = ", ".join(f'"{name}"' for name in names) + rows: list[dict[str, object]] = [] + for raw in backup.execute(f'select {quoted} from "{table.name}"'): + rows.append({name: _coerce(table.columns[name], raw[index]) for index, name in enumerate(names)}) + return rows + + +def _primary_key(table: Table) -> Column[Any]: + columns = list(table.primary_key.columns) + if len(columns) != 1: + message = f"{table.name} does not have a single-column primary key" + raise RuntimeError(message) + return columns[0] + + +def _existing_keys(connection: Connection, table: Table) -> set[object]: + key = _primary_key(table) + return set(connection.execute(select(key)).scalars().all()) + + +def _chunked(rows: Sequence[dict[str, object]], size: int = 200) -> Iterator[Sequence[dict[str, object]]]: + for start in range(0, len(rows), size): + yield rows[start : start + size] + + +def _verify_artifacts(settings: Settings) -> int: + """Re-hash every migrated artifact through the service's own verifier.""" + from transcription.db.models import ProcessingArtifact + from transcription.services.sources import SourceService + + engine = create_engine(_sync_url(settings)) + with engine.connect() as connection: + rows = connection.execute(select(SQLModel.metadata.tables["processing_artifact"])).mappings().all() + engine.dispose() + + service = SourceService(settings=settings) + artifacts = [ProcessingArtifact(**dict(row)) for row in rows] + # Reuses the application's own integrity check so the migration cannot + # disagree with what the running app considers a valid artifact. + service._verify_artifacts_integrity(artifacts) + return len(artifacts) + + +def _sync_url(settings: Settings) -> str: + """Return the target database URL with any async driver stripped.""" + url = get_database_url(settings) + return url.replace("+aiosqlite", "").replace("+asyncpg", "").replace("+psycopg", "") + + +def _preflight(backup: sqlite3.Connection, *, strict: bool) -> None: + actual = { + name: backup.execute(f'select count(*) from "{name}"').fetchone()[0] + for name in EXPECTED_SOURCE_COUNTS + } + mismatched = { + name: (count, EXPECTED_SOURCE_COUNTS[name]) + for name, count in actual.items() + if count != EXPECTED_SOURCE_COUNTS[name] + } + if not mismatched: + return + detail = ", ".join(f"{name}: found {found}, expected {want}" for name, (found, want) in sorted(mismatched.items())) + message = f"Backup row counts do not match the recorded V4.5 snapshot ({detail})" + if strict: + raise RuntimeError(message) + print(f"WARNING: {message}", file=sys.stderr) + + +def _copy_tables( + connection: Connection, + payload: dict[str, list[dict[str, object]]], + *, + dry_run: bool, +) -> tuple[int, dict[object, object]]: + """Insert every missing row, deferring the cyclic foreign key column.""" + deferred: dict[object, object] = {} + inserted_total = 0 + + for table in SQLModel.metadata.sorted_tables: + rows = payload[table.name] + existing = set() if dry_run else _existing_keys(connection, table) + key_name = _primary_key(table).name + + pending = [row for row in rows if row[key_name] not in existing] + + if table.name == DEFERRED_TABLE: + for row in pending: + selection = row[DEFERRED_COLUMN] + if selection is not None: + deferred[row[key_name]] = selection + row[DEFERRED_COLUMN] = None + + if pending and not dry_run: + for chunk in _chunked(pending): + connection.execute(insert(table), list(chunk)) + + inserted_total += len(pending) + print(f" {table.name:24} insert={len(pending):<5} skip={len(rows) - len(pending)}") + + return inserted_total, deferred + + +def _replay_deferred(connection: Connection, deferred: dict[object, object], *, dry_run: bool) -> None: + """Restore the preferred-attempt selections held back by the FK cycle.""" + if not deferred: + return + print(f" replaying {len(deferred)} deferred {DEFERRED_TABLE}.{DEFERRED_COLUMN} selection(s)") + if dry_run: + return + source = SQLModel.metadata.tables[DEFERRED_TABLE] + key = _primary_key(source) + for source_id, attempt_id in deferred.items(): + connection.execute(update(source).where(key == source_id).values({DEFERRED_COLUMN: attempt_id})) + + +def _report_counts(connection: Connection) -> None: + print("\nPost-migration row counts:") + for table in SQLModel.metadata.sorted_tables: + actual = len(connection.execute(select(_primary_key(table))).all()) + expected = EXPECTED_SOURCE_COUNTS.get(table.name) + flag = "" if expected is None or actual == expected else f" <-- expected {expected}" + print(f" {table.name:24} {actual}{flag}") + + +def _load_payload(backup_path: Path, *, strict_counts: bool) -> dict[str, list[dict[str, object]]]: + if not backup_path.is_file(): + message = f"Backup database not found: {backup_path}" + raise FileNotFoundError(message) + backup = sqlite3.connect(f"file:{backup_path}?mode=ro", uri=True) + try: + _preflight(backup, strict=strict_counts) + return {table.name: _read_table(backup, table) for table in SQLModel.metadata.sorted_tables} + finally: + backup.close() + + +def migrate(*, backup_path: Path, settings: Settings, dry_run: bool, strict_counts: bool) -> int: + """Copy every row from the V4.5 backup into the re-leveled schema.""" + payload = _load_payload(backup_path, strict_counts=strict_counts) + + engine = create_engine(_sync_url(settings)) + try: + if not sqlalchemy_inspect(engine).has_table("document"): + print("Target schema is empty; creating it from the current metadata.") + if not dry_run: + SQLModel.metadata.create_all(engine) + + with engine.begin() as connection: + inserted_total, deferred = _copy_tables(connection, payload, dry_run=dry_run) + _replay_deferred(connection, deferred, dry_run=dry_run) + + if dry_run: + print("\nDry run: no rows were written.") + return inserted_total + + with engine.connect() as connection: + _report_counts(connection) + finally: + engine.dispose() + + return inserted_total + + +def main(argv: Sequence[str] | None = None) -> int: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--backup", type=Path, default=DEFAULT_BACKUP, help="V4.5 database to read from") + parser.add_argument("--dry-run", action="store_true", help="Report what would be copied without writing") + parser.add_argument( + "--allow-count-mismatch", + action="store_true", + help="Warn instead of aborting when the backup row counts differ from the recorded snapshot", + ) + parser.add_argument( + "--verify-artifacts", + action="store_true", + help="Re-hash every migrated processing artifact after the copy", + ) + args = parser.parse_args(argv) + + settings = get_settings() + print(f"Source: {args.backup}") + print(f"Target: {_sync_url(settings)}\n") + + inserted = migrate( + backup_path=args.backup, + settings=settings, + dry_run=args.dry_run, + strict_counts=not args.allow_count_mismatch, + ) + + if args.verify_artifacts and not args.dry_run: + verified = _verify_artifacts(settings) + print(f"\nArtifact integrity verified for {verified} artifact(s).") + + print(f"\nDone. {inserted} row(s) inserted.") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main())