fix: 把同一张表的多个列放在一条 ALTER TABLE 里

This commit is contained in:
fawney19
2026-03-08 17:56:03 +08:00
parent 48d13762d9
commit 2d9158b321
@@ -117,21 +117,52 @@ def _is_numeric_type(table_name: str, column_name: str) -> bool:
return data_type == "numeric" return data_type == "numeric"
def upgrade() -> None: def _type_spec(col: str) -> str:
# -- 1. cost fields: Float -> Numeric """Return the SQL type literal for a given column name."""
for table, col, _nullable, _default in _COST_COLUMNS: return "NUMERIC(10,6)" if col == "rate_multiplier" else "NUMERIC(20,8)"
def _batch_alter_type(
columns: list[tuple[str, str, bool, str | None]],
cast_suffix: str,
type_fn=None,
) -> None:
"""Group columns by table and issue ONE ALTER TABLE per table.
This avoids rewriting the same table N times (once per column).
"""
from collections import defaultdict
by_table: dict[str, list[tuple[str, str, bool, str | None]]] = defaultdict(list)
for table, col, nullable, default in columns:
if not _column_exists(table, col): if not _column_exists(table, col):
continue continue
if _is_numeric_type(table, col): by_table[table].append((table, col, nullable, default))
bind = op.get_bind()
for table, cols in by_table.items():
# Build a single ALTER TABLE with multiple ALTER COLUMN clauses
parts: list[str] = []
for _t, col, _nullable, _default in cols:
target_type = type_fn(col) if type_fn else _type_spec(col)
parts.append(
f"ALTER COLUMN {col} TYPE {target_type} USING {col}::{cast_suffix}"
)
if not parts:
continue continue
target_type = _RATE_MULTIPLIER_TYPE if col == "rate_multiplier" else _COST_TYPE sql = f"ALTER TABLE {table} " + ", ".join(parts)
op.alter_column( bind.execute(sa.text(sql))
table,
col,
type_=target_type, def upgrade() -> None:
existing_type=sa.Float(), # -- 1. cost fields: Float -> Numeric (batched per table)
postgresql_using=f"{col}::numeric", # Filter out columns that are already numeric
) cols_to_convert = [
(t, c, n, d)
for t, c, n, d in _COST_COLUMNS
if _column_exists(t, c) and not _is_numeric_type(t, c)
]
_batch_alter_type(cols_to_convert, cast_suffix="numeric", type_fn=_type_spec)
# -- 2. provider_api_keys composite index # -- 2. provider_api_keys composite index
if not _index_exists("idx_provider_api_keys_provider_active"): if not _index_exists("idx_provider_api_keys_provider_active"):
@@ -150,16 +181,14 @@ def downgrade() -> None:
table_name="provider_api_keys", table_name="provider_api_keys",
) )
# -- 1. Numeric -> Float # -- 1. Numeric -> Float (batched per table)
for table, col, _nullable, _default in _COST_COLUMNS: cols_to_revert = [
if not _column_exists(table, col): (t, c, n, d)
continue for t, c, n, d in _COST_COLUMNS
if not _is_numeric_type(table, col): if _column_exists(t, c) and _is_numeric_type(t, c)
continue ]
op.alter_column( _batch_alter_type(
table, cols_to_revert,
col, cast_suffix="double precision",
type_=sa.Float(), type_fn=lambda _col: "DOUBLE PRECISION",
existing_type=_COST_TYPE if col != "rate_multiplier" else _RATE_MULTIPLIER_TYPE, )
postgresql_using=f"{col}::double precision",
)