"""repair tenant identity lookup indexes Revision ID: 20260718_0029 Revises: 20260717_0028 Create Date: 2026-07-18 """ from __future__ import annotations from collections.abc import Sequence import sqlalchemy as sa from alembic import op revision: str = "20260718_0029" down_revision: str | None = "20260717_0028" branch_labels: str | Sequence[str] | None = None depends_on: str | Sequence[str] | None = None class _TargetIndex: __slots__ = ( "table_name", "column_name", "index_name", "tenant_unique_constraint", ) def __init__( self, *, table_name: str, column_name: str, index_name: str, tenant_unique_constraint: str, ) -> None: self.table_name = table_name self.column_name = column_name self.index_name = index_name self.tenant_unique_constraint = tenant_unique_constraint _TARGET_INDEXES = ( _TargetIndex( table_name="organization_units", column_name="unit_code", index_name="ix_organization_units_unit_code", tenant_unique_constraint="uq_organization_units_tenant_code", ), _TargetIndex( table_name="employees", column_name="employee_no", index_name="ix_employees_employee_no", tenant_unique_constraint="uq_employees_tenant_employee_no", ), _TargetIndex( table_name="employees", column_name="email", index_name="ix_employees_email", tenant_unique_constraint="uq_employees_tenant_email", ), ) def _require_postgresql() -> None: dialect_name = op.get_bind().dialect.name if dialect_name != "postgresql": raise RuntimeError( "20260718_0029 only supports PostgreSQL; " f"refusing to mutate {dialect_name} without transactional index DDL" ) def _inspector() -> sa.Inspector: return sa.inspect(op.get_bind()) def _assert_tenant_unique_constraint(target: _TargetIndex) -> None: constraints = { str(item.get("name") or ""): tuple(item.get("column_names") or ()) for item in _inspector().get_unique_constraints(target.table_name) } expected_columns = ("tenant_id", target.column_name) actual_columns = constraints.get(target.tenant_unique_constraint) if actual_columns != expected_columns: raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"{target.tenant_unique_constraint} must cover {expected_columns}, " f"found {actual_columns!r}" ) def _validated_present_targets() -> tuple[_TargetIndex, ...]: inspector = _inspector() present_targets: list[_TargetIndex] = [] for target in _TARGET_INDEXES: if not inspector.has_table(target.table_name): continue columns = { str(item.get("name") or ""): item for item in inspector.get_columns(target.table_name) } required_columns = {"tenant_id", target.column_name} missing_columns = required_columns - columns.keys() if missing_columns: raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"{target.table_name} is missing columns " f"{', '.join(sorted(missing_columns))}" ) if bool(columns["tenant_id"].get("nullable", True)): raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"{target.table_name}.tenant_id must be non-nullable" ) _assert_tenant_unique_constraint(target) present_targets.append(target) return tuple(present_targets) def _index_uniqueness(target: _TargetIndex) -> bool | None: indexes = [ item for item in _inspector().get_indexes(target.table_name) if str(item.get("name") or "") == target.index_name ] if not indexes: return None if len(indexes) != 1: raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"found multiple indexes named {target.index_name}" ) index = indexes[0] actual_columns = tuple(index.get("column_names") or ()) if actual_columns != (target.column_name,): raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"{target.index_name} must cover only {target.column_name}, " f"found {actual_columns!r}" ) if index.get("duplicates_constraint"): raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"{target.index_name} backs a constraint and is not safe to replace" ) dialect_options = dict(index.get("dialect_options") or {}) predicate = dialect_options.get("postgresql_where") index_method = dialect_options.get("postgresql_using") operator_classes = dialect_options.get("postgresql_ops") or {} included_columns = ( index.get("include_columns") or dialect_options.get("postgresql_include") or () ) column_sorting = index.get("column_sorting") or {} if ( predicate is not None or included_columns or index_method not in (None, "btree") or operator_classes or column_sorting ): raise RuntimeError( "cannot repair tenant identity lookup indexes: " f"{target.index_name} is not a plain single-column btree index" ) return bool(index.get("unique", False)) def _replace_index(target: _TargetIndex, *, unique: bool) -> None: current_uniqueness = _index_uniqueness(target) if current_uniqueness == unique: return if current_uniqueness is not None: op.drop_index(target.index_name, table_name=target.table_name) op.create_index( target.index_name, target.table_name, [target.column_name], unique=unique, ) if _index_uniqueness(target) != unique: raise RuntimeError( "tenant identity lookup index replacement did not reach the expected state: " f"{target.index_name} unique={unique}" ) def _cross_tenant_duplicate_count(target: _TargetIndex) -> int: table = sa.table( target.table_name, sa.column("tenant_id"), sa.column(target.column_name), ) value_column = table.c[target.column_name] duplicate_values = ( sa.select(value_column) .where(value_column.is_not(None)) .group_by(value_column) .having(sa.func.count(sa.distinct(table.c.tenant_id)) > 1) .subquery() ) return int( op.get_bind().scalar( sa.select(sa.func.count()).select_from(duplicate_values) ) or 0 ) def upgrade() -> None: _require_postgresql() targets = _validated_present_targets() # 先校验所有旧索引形态,防止修到一半才发现同名索引承载了其他用途。 for target in targets: _index_uniqueness(target) for target in targets: _replace_index(target, unique=False) for target in targets: _assert_tenant_unique_constraint(target) def downgrade() -> None: _require_postgresql() targets = _validated_present_targets() # downgrade 会恢复旧的全局唯一索引;必须在任何 DDL 前一次性排除跨租户重复。 for target in targets: _index_uniqueness(target) violations = { f"{target.table_name}.{target.column_name}": duplicate_count for target in targets if (duplicate_count := _cross_tenant_duplicate_count(target)) > 0 } if violations: details = ", ".join( f"{name}={count}" for name, count in sorted(violations.items()) ) raise RuntimeError( "cannot downgrade tenant identity lookup indexes: " f"cross-tenant duplicate values exist ({details})" ) for target in targets: _replace_index(target, unique=True) for target in targets: _assert_tenant_unique_constraint(target)