232 lines
8.0 KiB
Python
232 lines
8.0 KiB
Python
"""将本项目 SQLite 数据安全迁移到 MySQL。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import json
|
||
import os
|
||
import subprocess
|
||
import sys
|
||
from collections.abc import Iterable
|
||
from pathlib import Path
|
||
|
||
from sqlalchemy import JSON, MetaData, create_engine, func, inspect, select
|
||
from sqlalchemy.engine import Connection, Engine
|
||
from sqlalchemy.schema import Table
|
||
|
||
|
||
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
||
DEFAULT_SOURCE = PROJECT_ROOT / "data" / "web.db"
|
||
IGNORED_SOURCE_TABLES = {"alembic_version"}
|
||
|
||
if str(PROJECT_ROOT) not in sys.path:
|
||
sys.path.insert(0, str(PROJECT_ROOT))
|
||
|
||
|
||
def parse_args() -> argparse.Namespace:
|
||
"""解析迁移参数。"""
|
||
parser = argparse.ArgumentParser(description="迁移 SQLite 数据到 MySQL")
|
||
parser.add_argument(
|
||
"--source",
|
||
type=Path,
|
||
default=DEFAULT_SOURCE,
|
||
help=f"SQLite 数据库路径,默认:{DEFAULT_SOURCE}",
|
||
)
|
||
parser.add_argument(
|
||
"--target-url",
|
||
help="MySQL SQLAlchemy 连接串,例如 mysql+pymysql://user:password@host:3306/database",
|
||
)
|
||
parser.add_argument(
|
||
"--batch-size",
|
||
type=int,
|
||
default=1000,
|
||
help="每批写入行数,默认 1000",
|
||
)
|
||
args = parser.parse_args()
|
||
if args.batch_size < 1:
|
||
parser.error("--batch-size 必须大于 0")
|
||
return args
|
||
|
||
|
||
def _count_rows(connection: Connection, table: Table) -> int:
|
||
return connection.scalar(select(func.count()).select_from(table)) or 0
|
||
|
||
|
||
def _chunks(rows: Iterable[dict], batch_size: int) -> Iterable[list[dict]]:
|
||
"""将行数据分批,控制单次 SQL 的内存与事务体积。"""
|
||
batch: list[dict] = []
|
||
for row in rows:
|
||
batch.append(row)
|
||
if len(batch) >= batch_size:
|
||
yield batch
|
||
batch = []
|
||
if batch:
|
||
yield batch
|
||
|
||
|
||
def _normalize_value(value: object, target_column) -> object:
|
||
"""将 SQLite 中以文本保存的 JSON 转回结构化值。"""
|
||
if value is None or not isinstance(target_column.type, JSON):
|
||
return value
|
||
if not isinstance(value, str):
|
||
return value
|
||
try:
|
||
return json.loads(value)
|
||
except json.JSONDecodeError as exc:
|
||
raise ValueError(
|
||
f"列 {target_column.name} 存在非法 JSON:{value[:120]!r}"
|
||
) from exc
|
||
|
||
|
||
def _source_rows(
|
||
source_connection: Connection, source_table: Table, target_table: Table
|
||
) -> Iterable[dict]:
|
||
columns = [column.name for column in target_table.columns]
|
||
missing_columns = [column for column in columns if column not in source_table.c]
|
||
if missing_columns:
|
||
raise RuntimeError(
|
||
f"源表 {source_table.name} 缺少列:{', '.join(missing_columns)},"
|
||
"请先将 SQLite 升级到当前版本后再迁移"
|
||
)
|
||
|
||
primary_key_columns = list(source_table.primary_key.columns)
|
||
statement = select(*(source_table.c[name] for name in columns))
|
||
if primary_key_columns:
|
||
statement = statement.order_by(
|
||
*(column.asc() for column in primary_key_columns)
|
||
)
|
||
|
||
for row in source_connection.execute(statement).mappings():
|
||
yield {
|
||
column_name: _normalize_value(row[column_name], target_table.c[column_name])
|
||
for column_name in columns
|
||
}
|
||
|
||
|
||
def _ensure_target_is_empty(target_engine: Engine, target_tables: list[Table]) -> None:
|
||
"""拒绝向已有业务数据的 MySQL 写入,防止误覆盖。"""
|
||
with target_engine.connect() as connection:
|
||
occupied = [
|
||
table.name for table in target_tables if _count_rows(connection, table) > 0
|
||
]
|
||
if occupied:
|
||
raise RuntimeError(
|
||
f"目标 MySQL 已存在业务数据,拒绝迁移:{', '.join(occupied)}"
|
||
)
|
||
|
||
|
||
def _upgrade_source_sqlite(source_path: Path) -> None:
|
||
"""先将源 SQLite 升级到当前 Alembic 版本,补齐历史表字段。"""
|
||
print("正在升级源 SQLite 表结构...")
|
||
source_env = os.environ.copy()
|
||
source_env["DATABASE_URL"] = f"sqlite:///{source_path}"
|
||
subprocess.run(
|
||
[
|
||
sys.executable,
|
||
"-c",
|
||
"from web.backend.database import run_migrations; run_migrations()",
|
||
],
|
||
cwd=PROJECT_ROOT,
|
||
env=source_env,
|
||
check=True,
|
||
)
|
||
|
||
|
||
def main() -> int:
|
||
"""执行迁移、计数核验并返回进程退出码。"""
|
||
args = parse_args()
|
||
source_path = args.source.expanduser().resolve()
|
||
target_url = (args.target_url or "").strip()
|
||
|
||
if not source_path.is_file():
|
||
raise FileNotFoundError(f"未找到 SQLite 数据库:{source_path}")
|
||
|
||
_upgrade_source_sqlite(source_path)
|
||
|
||
# 必须在导入数据库模块前设置,Alembic 环境才会使用迁移目标库。
|
||
if target_url:
|
||
os.environ["DATABASE_URL"] = target_url
|
||
from web.backend import models # noqa: F401
|
||
from web.backend.database import Base, DATABASE_URL, run_migrations
|
||
|
||
target_url = target_url or DATABASE_URL
|
||
if not target_url.startswith("mysql+"):
|
||
raise ValueError(
|
||
"目标库必须是 MySQL;可传 --target-url,或设置 DB_HOST/DB_* 环境变量"
|
||
)
|
||
|
||
print("正在初始化目标 MySQL 表结构...")
|
||
run_migrations()
|
||
|
||
target_engine = create_engine(target_url, pool_pre_ping=True)
|
||
source_engine = create_engine(f"sqlite:///{source_path}")
|
||
target_tables = list(Base.metadata.sorted_tables)
|
||
target_table_names = {table.name for table in target_tables}
|
||
|
||
try:
|
||
source_table_names = set(inspect(source_engine).get_table_names())
|
||
unexpected_tables = (
|
||
source_table_names - target_table_names - IGNORED_SOURCE_TABLES
|
||
)
|
||
if unexpected_tables:
|
||
raise RuntimeError(
|
||
"源 SQLite 存在当前程序无法识别的表:"
|
||
f"{', '.join(sorted(unexpected_tables))},为避免遗漏已停止迁移"
|
||
)
|
||
|
||
source_metadata = MetaData()
|
||
source_metadata.reflect(
|
||
bind=source_engine,
|
||
only=sorted(source_table_names & target_table_names),
|
||
)
|
||
tables_to_copy = [
|
||
table for table in target_tables if table.name in source_metadata.tables
|
||
]
|
||
|
||
_ensure_target_is_empty(target_engine, tables_to_copy)
|
||
|
||
with (
|
||
source_engine.connect() as source_connection,
|
||
target_engine.begin() as target_connection,
|
||
):
|
||
for target_table in tables_to_copy:
|
||
source_table = source_metadata.tables[target_table.name]
|
||
source_count = _count_rows(source_connection, source_table)
|
||
if not source_count:
|
||
print(f"{target_table.name}: 0 行,跳过")
|
||
continue
|
||
|
||
print(f"{target_table.name}: 正在迁移 {source_count} 行...")
|
||
rows = _source_rows(source_connection, source_table, target_table)
|
||
for batch in _chunks(rows, args.batch_size):
|
||
target_connection.execute(target_table.insert(), batch)
|
||
|
||
with (
|
||
source_engine.connect() as source_connection,
|
||
target_engine.connect() as target_connection,
|
||
):
|
||
for target_table in tables_to_copy:
|
||
source_table = source_metadata.tables[target_table.name]
|
||
source_count = _count_rows(source_connection, source_table)
|
||
target_count = _count_rows(target_connection, target_table)
|
||
if source_count != target_count:
|
||
raise RuntimeError(
|
||
f"{target_table.name} 行数核验失败:"
|
||
f"SQLite={source_count},MySQL={target_count}"
|
||
)
|
||
print(f"{target_table.name}: 核验通过({target_count} 行)")
|
||
finally:
|
||
source_engine.dispose()
|
||
target_engine.dispose()
|
||
|
||
print("迁移完成,SQLite 业务数据未修改,表结构已升级到当前版本。")
|
||
return 0
|
||
|
||
|
||
if __name__ == "__main__":
|
||
try:
|
||
sys.exit(main())
|
||
except Exception as exc:
|
||
print(f"迁移失败:{exc}", file=sys.stderr)
|
||
sys.exit(1)
|