Files
live-hub-py/scripts/migrate_sqlite_to_mysql.py
T

212 lines
7.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""将本项目 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)