Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,9 @@ jobs:
./scripts/check-rust-file-size
fi

- name: Migration version guard
run: ./scripts/check-migration-versions

- name: Install native build dependencies
run: |
sudo apt-get update
Expand Down
91 changes: 91 additions & 0 deletions scripts/check-migration-versions
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
#!/usr/bin/env python3
"""Check SQLx migration filenames for well-formed, unique versions."""
from __future__ import annotations

import argparse
import re
import sys
from collections import defaultdict
from pathlib import Path

REPO_ROOT = Path(__file__).resolve().parent.parent
DEFAULT_MIGRATIONS_DIR = REPO_ROOT / "migrations"
MIGRATION_FILENAME = re.compile(r"^(?P<version>[0-9]+)_.+\.sql$")


def display_path(path: Path) -> str:
try:
return path.relative_to(REPO_ROOT).as_posix()
except ValueError:
return path.as_posix()


def check(migrations_dir: Path) -> int:
if not migrations_dir.is_dir():
print(
f"Migration version check failed: {display_path(migrations_dir)} "
"is not a directory.",
file=sys.stderr,
)
return 1

migration_files = sorted(
path for path in migrations_dir.glob("*.sql") if path.is_file()
)
invalid_filenames: list[Path] = []
files_by_version: dict[int, list[Path]] = defaultdict(list)

for path in migration_files:
match = MIGRATION_FILENAME.fullmatch(path.name)
if match is None:
invalid_filenames.append(path)
continue
files_by_version[int(match.group("version"))].append(path)

collisions = {
version: paths
for version, paths in files_by_version.items()
if len(paths) > 1
}
if invalid_filenames or collisions:
print("Migration version check failed:", file=sys.stderr)
if invalid_filenames:
print(
" Invalid filenames (expected <digits>_<name>.sql):",
file=sys.stderr,
)
for path in invalid_filenames:
print(f" - {display_path(path)}", file=sys.stderr)
if collisions:
print(" Duplicate migration versions:", file=sys.stderr)
for version, paths in sorted(collisions.items()):
print(f" {version}:", file=sys.stderr)
for path in paths:
print(f" - {display_path(path)}", file=sys.stderr)
return 1

print(
f"Migration version check OK: {len(migration_files)} migration files "
"have unique versions."
)
return 0


def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--migrations-dir",
type=Path,
default=DEFAULT_MIGRATIONS_DIR,
help="migration directory (default: migrations/)",
)
args = parser.parse_args()

migrations_dir = args.migrations_dir
if not migrations_dir.is_absolute():
migrations_dir = REPO_ROOT / migrations_dir
return check(migrations_dir)


if __name__ == "__main__":
sys.exit(main())
Loading