from importlib.util import module_from_spec, spec_from_file_location from pathlib import Path import sqlalchemy as sa VERSIONS_DIR = Path(__file__).parents[1] / "alembic" / "versions" def load_migration(): path = VERSIONS_DIR / "0034_media_assets.py" spec = spec_from_file_location("media_asset_migration", path) assert spec and spec.loader migration = module_from_spec(spec) spec.loader.exec_module(migration) return migration class Inspector: def __init__(self, tables: set[str]): self.tables = tables def get_table_names(self) -> list[str]: return sorted(self.tables) def test_media_asset_migration_creates_the_upload_table_with_scope_columns(monkeypatch): migration = load_migration() created: dict[str, object] = {} indexes: list[tuple[object, ...]] = [] monkeypatch.setattr(migration, "inspect", lambda _bind: Inspector({"AdminUser", "AdminDepartment"})) monkeypatch.setattr(migration.op, "get_bind", lambda: object()) monkeypatch.setattr( migration.op, "create_table", lambda *args, **kwargs: created.update(args=args, kwargs=kwargs), ) monkeypatch.setattr( migration.op, "create_index", lambda *args, **kwargs: indexes.append((*args, *kwargs.values())), ) migration.upgrade() assert migration.down_revision == "0033_remove_site_versions" assert created["args"][0] == "MediaAsset" columns = {column.name: column for column in created["args"][1:] if isinstance(column, sa.Column)} assert set(columns) == { "id", "url", "name", "mimeType", "sizeBytes", "group", "deptId", "createdById", "createdAt", "updatedAt", } assert columns["url"].unique assert columns["deptId"].nullable is False assert columns["deptId"].server_default.arg.text.strip("'") == migration.DEFAULT_DEPT_ID assert columns["createdById"].nullable is True assert {index[0] for index in indexes} == { "ix_MediaAsset_deptId", "ix_MediaAsset_createdById", } constraints = [constraint for constraint in created["args"][1:] if isinstance(constraint, sa.ForeignKeyConstraint)] primary_keys = [constraint for constraint in created["args"][1:] if isinstance(constraint, sa.PrimaryKeyConstraint)] assert primary_keys[0]._pending_colargs == ["id"] assert { (constraint.name, tuple(constraint.column_keys), tuple(constraint.elements[0].target_fullname.split("."))) for constraint in constraints } == { ("fk_MediaAsset_deptId", ("deptId",), ("AdminDepartment", "id")), ("fk_MediaAsset_createdById", ("createdById",), ("AdminUser", "id")), } def test_media_asset_migration_is_safe_when_table_already_exists(monkeypatch): migration = load_migration() calls: list[tuple[object, ...]] = [] monkeypatch.setattr(migration, "inspect", lambda _bind: Inspector({"MediaAsset"})) monkeypatch.setattr(migration.op, "get_bind", lambda: object()) monkeypatch.setattr(migration.op, "create_table", lambda *args, **_kwargs: calls.append(args)) monkeypatch.setattr(migration.op, "create_index", lambda *args, **_kwargs: calls.append(args)) migration.upgrade() assert calls == []