diff --git a/scripts/check_migration_safety.py b/scripts/check_migration_safety.py new file mode 100644 index 000000000..6ec4b784d --- /dev/null +++ b/scripts/check_migration_safety.py @@ -0,0 +1,250 @@ +#!/usr/bin/env python3 +""" +数据库迁移破坏性变更安全检查 + +只检查 Alembic 迁移文件的 upgrade 函数中是否包含破坏性操作: +- DROP TABLE +- ALTER TABLE ... DROP COLUMN +- 列类型变更(可能导致数据丢失) +- NOT NULL 约束新增(无默认值时) +- 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" + +# 高风险模式:直接导致数据丢失(只在 upgrade 中检查) +HIGH_RISK_PATTERNS = [ + (r"\bop\.drop_table\(", "op.drop_table() - 删除表,数据永久丢失"), + (r"\bop\.drop_column\(", "op.drop_column() - 删除列,数据永久丢失"), +] + +# 中风险模式:可能导致数据丢失或兼容性问题 +MEDIUM_RISK_PATTERNS = [ + (r"op\.alter_column\([^)]*nullable\s*=\s*False", "新增 NOT NULL 约束 - 旧数据可能为空导致迁移失败"), + (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 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}") + + 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() + + 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()) +