Files
X-Financial/server/alembic/versions/20260718_0029_tenant_identity_lookup_indexes.py

250 lines
7.9 KiB
Python
Raw Normal View History

"""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)