133 lines
3.0 KiB
Python
133 lines
3.0 KiB
Python
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import sqlite3
|
|
from pathlib import Path
|
|
|
|
from .connection import DEFAULT_DATABASE_PATH, connect_database
|
|
|
|
|
|
MIGRATIONS_DIR = Path(__file__).resolve().parents[1] / "migrations"
|
|
|
|
|
|
def _migration_checksum(path: Path) -> str:
|
|
digest = hashlib.sha256()
|
|
with path.open("rb") as fh:
|
|
for chunk in iter(lambda: fh.read(1024 * 1024), b""):
|
|
digest.update(chunk)
|
|
return digest.hexdigest()
|
|
|
|
|
|
def _ensure_migrations_table(conn: sqlite3.Connection) -> None:
|
|
conn.execute(
|
|
"""
|
|
CREATE TABLE IF NOT EXISTS cracklab_migrations (
|
|
id TEXT PRIMARY KEY,
|
|
checksum TEXT NOT NULL,
|
|
applied_at TEXT NOT NULL
|
|
)
|
|
"""
|
|
)
|
|
|
|
|
|
def _load_applied_migrations(
|
|
conn: sqlite3.Connection,
|
|
) -> dict[str, tuple[str, str]]:
|
|
rows = conn.execute(
|
|
"""
|
|
SELECT id, checksum, applied_at
|
|
FROM cracklab_migrations
|
|
"""
|
|
).fetchall()
|
|
|
|
return {
|
|
migration_id: (checksum, applied_at)
|
|
for migration_id, checksum, applied_at in rows
|
|
}
|
|
|
|
|
|
def _migration_files(migrations_dir: Path) -> list[Path]:
|
|
return sorted(
|
|
path
|
|
for path in migrations_dir.glob("*.sql")
|
|
if path.is_file()
|
|
)
|
|
|
|
|
|
def migrate_database(
|
|
path: Path | None = None,
|
|
migrations_dir: Path | None = None,
|
|
) -> list[str]:
|
|
conn = connect_database(path)
|
|
|
|
try:
|
|
_ensure_migrations_table(conn)
|
|
conn.commit()
|
|
|
|
applied = _load_applied_migrations(conn)
|
|
migration_directory = migrations_dir or MIGRATIONS_DIR
|
|
migrations = _migration_files(migration_directory)
|
|
newly_applied: list[str] = []
|
|
|
|
for migration_path in migrations:
|
|
migration_id = migration_path.name
|
|
checksum = _migration_checksum(migration_path)
|
|
|
|
if migration_id in applied:
|
|
stored_checksum, _ = applied[migration_id]
|
|
|
|
if stored_checksum != checksum:
|
|
raise RuntimeError(
|
|
f"Migration checksum mismatch: {migration_id}"
|
|
)
|
|
|
|
continue
|
|
|
|
sql = migration_path.read_text(encoding="utf-8")
|
|
|
|
transaction_sql = f"""
|
|
BEGIN;
|
|
{sql}
|
|
INSERT INTO cracklab_migrations (
|
|
id,
|
|
checksum,
|
|
applied_at
|
|
)
|
|
VALUES (
|
|
'{migration_id.replace("'", "''")}',
|
|
'{checksum}',
|
|
strftime('%Y-%m-%dT%H:%M:%fZ', 'now')
|
|
);
|
|
COMMIT;
|
|
"""
|
|
|
|
try:
|
|
conn.executescript(transaction_sql)
|
|
except Exception:
|
|
conn.rollback()
|
|
raise
|
|
|
|
applied[migration_id] = (
|
|
checksum,
|
|
"",
|
|
)
|
|
newly_applied.append(migration_id)
|
|
|
|
return newly_applied
|
|
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def main() -> int:
|
|
applied = migrate_database(DEFAULT_DATABASE_PATH)
|
|
|
|
for migration_id in applied:
|
|
print(migration_id)
|
|
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|