#!/usr/bin/env python3 """ 数据库迁移破坏性变更安全检查 只检查 Alembic 迁移文件的 upgrade 函数中是否包含破坏性操作: - DROP TABLE / DROP COLUMN(高风险,直接阻断) - 不加默认值的 NOT NULL 列添加(中风险,告警) - 大表的 ALTER TABLE(超过10万行,中风险,告警) - 列类型变更(可能导致数据丢失,中风险) - RENAME TABLE / RENAME COLUMN(中风险) 忽略 downgrade 函数中的操作(那是回滚逻辑,正常的)。 使用方式: # 检查所有迁移(不推荐,会扫历史已执行的迁移) python3 scripts/check_migration_safety.py # 只检查与目标分支相比新增的迁移(推荐用于CI) python3 scripts/check_migration_safety.py --diff-against origin/main # 只检查指定版本之后的迁移 python3 scripts/check_migration_safety.py --since 030_xxx 退出码: 0 - 安全 / 只有非破坏性变更 1 - 检测到高风险破坏性变更 2 - 检测到中风险变更,需人工确认 """ from __future__ import annotations import argparse import os import re import subprocess import sys from pathlib import Path from typing import List, Tuple REPO_ROOT = Path(__file__).resolve().parents[1] ALEMBIC_VERSIONS_DIR = REPO_ROOT / "alembic" / "versions" # 已知可能的大表(超过10万行的表),ALTER TABLE 这些表需要告警 # 可通过环境变量 MIGRATION_LARGE_TABLES 覆盖,用逗号分隔 LARGE_TABLES = os.getenv( "MIGRATION_LARGE_TABLES", "generation_projects,generation_tasks,assets,upload_sessions,users", ).split(",") LARGE_TABLE_THRESHOLD = int(os.getenv("MIGRATION_LARGE_TABLE_THRESHOLD", "100000")) # 高风险模式:直接导致数据丢失(只在 upgrade 中检查) HIGH_RISK_PATTERNS = [ (r"\bop\.drop_table\(", "op.drop_table() - 删除表,数据永久丢失"), (r"\bop\.drop_column\(", "op.drop_column() - 删除列,数据永久丢失"), ] # 中风险模式:可能导致数据丢失或兼容性问题 # 注意:NOT NULL 检测需要额外判断是否有默认值,使用专用函数检测 MEDIUM_RISK_PATTERNS = [ (r"op\.alter_column\([^)]*type_\s*=", "列类型变更 - 可能导致数据截断或转换失败"), ( r"\bop\.rename_table\(", "op.rename_table() - 重命名表,可能导致依赖该表的代码报错", ), ( r"\bop\.rename_column\(", "op.rename_column() - 重命名列,可能导致依赖该列的代码报错", ), (r"\bop\.drop_index\(", "op.drop_index() - 删除索引,可能影响查询性能"), (r"\bop\.drop_constraint\(", "op.drop_constraint() - 删除约束,可能影响数据完整性"), ] # 安全模式:这些是安全的新增操作 SAFE_PATTERNS = [ (r"\bop\.create_table\(", "新建表"), (r"\bop\.add_column\(", "新增列"), (r"\bop\.create_index\(", "新建索引"), (r"\bop\.create_unique_constraint\(", "新建唯一约束"), (r"\bop\.create_foreign_key\(", "新建外键约束"), ] def extract_upgrade_content(content: str) -> str: """ 从迁移文件中提取 upgrade 函数的内容。 只检查 upgrade 中的操作,忽略 downgrade。 """ upgrade_match = re.search(r"def upgrade\b[^:]*:", content) if not upgrade_match: return "" upgrade_start = upgrade_match.end() # 找到下一个顶层 def(通常是 def downgrade)作为结束位置 rest = content[upgrade_start:] downgrade_match = re.search(r"\n\ndef\s+\w+\b", rest) if downgrade_match: upgrade_end = upgrade_start + downgrade_match.start() else: upgrade_end = len(content) return content[upgrade_start:upgrade_end] def extract_alter_column_calls(upgrade_content: str) -> List[str]: """ 从 upgrade 内容中提取所有 op.alter_column() 调用的完整参数字符串。 处理跨行的情况。 """ calls = [] pattern = r"op\.alter_column\(" pos = 0 while True: match = re.search(pattern, upgrade_content[pos:]) if not match: break call_start = pos + match.start() # 从 opening 括号之后开始计数,初始深度为 1 i = call_start + len("op.alter_column(") depth = 1 while i < len(upgrade_content) and depth > 0: if upgrade_content[i] == "(": depth += 1 elif upgrade_content[i] == ")": depth -= 1 i += 1 if depth == 0: # 找到了匹配的闭合括号,i 指向闭合括号之后 calls.append(upgrade_content[call_start : i - 1]) pos = i return calls def check_not_null_without_default(file_path: Path, upgrade_content: str) -> List[str]: """ 检测 NOT NULL 列添加但没有默认值的情况。 只有 nullable=False 且没有 server_default / default 时才告警。 """ issues = [] alter_calls = extract_alter_column_calls(upgrade_content) for call in alter_calls: # 检查是否设置了 nullable=False if not re.search(r"nullable\s*=\s*False", call): continue # 检查是否有默认值(server_default 或 default) has_default = bool(re.search(r"server_default\s*=", call) or re.search(r"\bdefault\s*=", call)) if not has_default: # 提取表名和列名 # op.alter_column('table_name', 'column_name', ...) table_col_match = re.search( r"op\.alter_column\(\s*['\"]([^'\"]+)['\"]\s*,\s*['\"]([^'\"]+)['\"]", call, ) if table_col_match: table = table_col_match.group(1) col = table_col_match.group(2) issues.append( f"{file_path.name}: 新增 NOT NULL 约束无默认值 " f"(表: {table}, 列: {col})- 旧数据为空时迁移失败" ) else: issues.append(f"{file_path.name}: 新增 NOT NULL 约束无默认值 " "- 旧数据为空时迁移失败") return issues def check_large_table_alter(file_path: Path, upgrade_content: str) -> List[str]: """ 检测大表的 ALTER TABLE 操作。 大表名单通过 LARGE_TABLES 配置。 """ issues = [] # 检测 alter_column 在大表上的操作 alter_calls = extract_alter_column_calls(upgrade_content) for call in alter_calls: table_col_match = re.search( r"op\.alter_column\(\s*['\"]([^'\"]+)['\"]", call, ) if not table_col_match: continue table = table_col_match.group(1) if table in LARGE_TABLES: # 提取操作类型 ops = [] if re.search(r"nullable\s*=", call): ops.append("修改nullable") if re.search(r"type_\s*=", call): ops.append("列类型变更") if re.search(r"server_default\s*=", call): ops.append("修改默认值") if not ops: ops.append("ALTER COLUMN") issues.append( f"{file_path.name}: 大表 ALTER TABLE " f"(表: {table}, 操作: {', '.join(ops)})" f" - 预计行数 > {LARGE_TABLE_THRESHOLD},可能导致长时间锁表" ) # 检测大表上的 drop_column drop_col_pattern = r"op\.drop_column\(\s*['\"]([^'\"]+)['\"]\s*,\s*['\"]([^'\"]+)['\"]" for match in re.finditer(drop_col_pattern, upgrade_content): table = match.group(1) col = match.group(2) if table in LARGE_TABLES: # 这个已经算高风险了,但额外加上大表提示 pass # drop_column 已在 HIGH_RISK 中覆盖 return issues def get_new_migrations_via_diff(diff_target: str) -> List[Path]: """ 通过 git diff 对比目标分支/commit,找出 alembic/versions/ 下新增的迁移文件。 只包含新增文件(A状态),不包含修改或删除的文件。 """ try: result = subprocess.run( [ "git", "diff", "--name-only", "--diff-filter=A", diff_target, "HEAD", "--", "alembic/versions/", ], cwd=str(REPO_ROOT), capture_output=True, text=True, check=True, ) files = [line.strip() for line in result.stdout.strip().split("\n") if line.strip()] return [REPO_ROOT / f for f in files] except subprocess.CalledProcessError as e: print(f"⚠️ git diff 失败({diff_target}):{e.stderr.strip()}") print(f" 降级为检查所有迁移文件") return sorted(ALEMBIC_VERSIONS_DIR.glob("*.py")) def find_new_migrations(since_revision: str | None = None, diff_against: str | None = None) -> List[Path]: """ 找出需要检查的迁移文件。 优先级:diff_against > since_revision > 全部 """ if diff_against: return get_new_migrations_via_diff(diff_against) all_migrations = sorted(ALEMBIC_VERSIONS_DIR.glob("*.py")) if not since_revision: return all_migrations result = [] found = False for m in all_migrations: if since_revision in m.name or since_revision in m.stem: found = True continue if found: result.append(m) return result if found else all_migrations def analyze_migration(file_path: Path) -> Tuple[List[str], List[str], List[str]]: """分析单个迁移文件 upgrade 部分的风险等级""" content = file_path.read_text() upgrade_content = extract_upgrade_content(content) if not upgrade_content: return [], [], [f"{file_path.name}: 未找到 upgrade 函数"] high_risks = [] medium_risks = [] safes = [] # 高风险检测 for pattern, desc in HIGH_RISK_PATTERNS: if re.search(pattern, upgrade_content): high_risks.append(f"{file_path.name}: {desc}") # 中风险模式检测(通用正则) for pattern, desc in MEDIUM_RISK_PATTERNS: if re.search(pattern, upgrade_content): medium_risks.append(f"{file_path.name}: {desc}") # NOT NULL 无默认值检测(专用精确检测) not_null_issues = check_not_null_without_default(file_path, upgrade_content) medium_risks.extend(not_null_issues) # 大表 ALTER TABLE 检测 large_table_issues = check_large_table_alter(file_path, upgrade_content) medium_risks.extend(large_table_issues) # 安全操作检测 for pattern, desc in SAFE_PATTERNS: if re.search(pattern, upgrade_content): safes.append(f"{file_path.name}: {desc}") return high_risks, medium_risks, safes def main() -> int: parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) parser.add_argument( "--since", default=os.getenv("MIGRATION_SINCE_REVISION"), help="只检查指定版本之后的迁移(如:030_xxx),不传则检查所有迁移", ) parser.add_argument( "--diff-against", default=os.getenv("MIGRATION_DIFF_AGAINST"), help="对比指定分支/commit,只检查新增的迁移文件(推荐用于CI,如 origin/main)", ) parser.add_argument( "--warn-only", action="store_true", help="只警告不失败(用于非强制门禁场景)", ) parser.add_argument( "--allow-medium-risk", action="store_true", help="允许中风险变更(只拦截高风险)", ) args = parser.parse_args() migrations = find_new_migrations(args.since, args.diff_against) if not migrations: print("✅ 未找到需要检查的新增迁移文件,跳过") return 0 print(f"🔍 正在检查 {len(migrations)} 个迁移文件的 upgrade 操作...") if args.diff_against: print(f" (对比基准:{args.diff_against},仅检查新增迁移)") print() print(f"📋 大表名单(ALTER TABLE 会额外告警): {', '.join(LARGE_TABLES)}") print(f"📊 大表阈值: > {LARGE_TABLE_THRESHOLD} 行") print() all_high = [] all_medium = [] all_safe = [] for m in migrations: high, medium, safe = analyze_migration(m) all_high.extend(high) all_medium.extend(medium) all_safe.extend(safe) if all_safe: print("✅ 安全变更:") for s in all_safe: print(f" - {s}") print() if all_medium: print("⚠️ 中风险变更(需人工确认):") for m_item in all_medium: print(f" - {m_item}") print() if all_high: print("❌ 高风险破坏性变更(禁止自动部署):") for h in all_high: print(f" - {h}") print() print("=" * 60) print(f"检查结果:{len(all_safe)} 项安全 / {len(all_medium)} 项中风险 / {len(all_high)} 项高风险") print() if all_high: print("❌ 检测到高风险破坏性变更,CI 检查失败!") print(" 如果确认这是预期操作,请在 MR/PR 中说明原因并获得审批。") if args.warn_only: return 0 return 1 if all_medium and not args.allow_medium_risk: print("⚠️ 检测到中风险变更,请人工确认后再部署。") if args.warn_only: return 0 print("(如需仅拦截高风险,可使用 --allow-medium-risk 参数)") return 2 print("✅ 未检测到破坏性变更") return 0 if __name__ == "__main__": sys.exit(main())