Files
dy/backend/models/db_transfer.py
T
2026-07-23 17:56:25 +08:00

324 lines
12 KiB
Python

"""从 SQLite 源库迁移数据到目标数据库。"""
from __future__ import annotations
from datetime import date, datetime
from pathlib import Path
from typing import Any
from sqlalchemy import MetaData, insert, inspect, text
from sqlalchemy.ext.asyncio import AsyncEngine
from models.database import Base
import models.models # noqa: F401 — 注册 ORM 表结构
from models.db_config import (
PASSWORD_PLACEHOLDER,
DatabaseConfig,
build_database_url,
create_database_engine,
mask_database_url,
read_database_config,
)
from models.db_migrate import (
migrate_accounts_quota_disabled,
migrate_accounts_table,
migrate_message_logs_table,
migrate_payment_orders_table,
migrate_rules_table,
migrate_users_table,
)
# 按外键依赖顺序导入
MIGRATION_TABLES: tuple[str, ...] = (
"users",
"email_verification_tokens",
"password_reset_tokens",
"app_config",
"accounts",
"rules",
"message_logs",
"received_message_logs",
"system_logs",
"payment_orders",
)
BATCH_SIZE = 400
def _parse_datetime(value: Any) -> Any:
if value is None or isinstance(value, (datetime, date)):
return value
if isinstance(value, str):
text_value = value.strip()
if not text_value:
return None
for fmt in ("%Y-%m-%d %H:%M:%S.%f", "%Y-%m-%d %H:%M:%S", "%Y-%m-%d"):
try:
return datetime.strptime(text_value, fmt)
except ValueError:
continue
try:
return datetime.fromisoformat(text_value.replace("Z", "+00:00"))
except ValueError:
return value
return value
def _normalize_rows(rows: list[dict[str, Any]], tbl) -> list[dict[str, Any]]:
if not rows:
return rows
col_types = {column.name: column.type for column in tbl.columns}
normalized: list[dict[str, Any]] = []
for row in rows:
item = dict(row)
for key, col_type in col_types.items():
if key not in item or item[key] is None:
continue
type_name = col_type.__class__.__name__.lower()
if "datetime" in type_name or "timestamp" in type_name:
item[key] = _parse_datetime(item[key])
elif "boolean" in type_name and not isinstance(item[key], bool):
if isinstance(item[key], (int, float)):
item[key] = bool(item[key])
elif isinstance(item[key], str):
item[key] = item[key].strip().lower() in ("1", "true", "t", "yes")
normalized.append(item)
return normalized
def resolve_sqlite_source_path(source_path: str | None = None) -> Path:
if source_path and str(source_path).strip():
path = Path(source_path.strip())
if not path.is_absolute():
from models.db_config import BACKEND_DIR
path = BACKEND_DIR / path
return path
from models.db_config import DEFAULT_SQLITE_PATH
return DEFAULT_SQLITE_PATH
def build_sqlite_config(path: Path) -> DatabaseConfig:
return DatabaseConfig(db_type="sqlite", db_path=str(path))
def _same_sqlite_target(source: Path, target: DatabaseConfig) -> bool:
if target.normalized_type() != "sqlite":
return False
try:
return source.resolve() == target.sqlite_path().resolve()
except OSError:
return str(source) == str(target.sqlite_path())
async def inspect_sqlite_source(source_path: str | None = None) -> dict[str, Any]:
path = resolve_sqlite_source_path(source_path)
if not path.exists():
return {
"source_path": str(path),
"exists": False,
"tables": {},
"total_rows": 0,
}
cfg = build_sqlite_config(path)
engine = create_database_engine(cfg)
tables: dict[str, int] = {}
total_rows = 0
try:
async with engine.connect() as conn:
for table in MIGRATION_TABLES:
if not await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
tables[table] = 0
continue
count = await conn.scalar(text(f"SELECT COUNT(*) FROM {table}"))
row_count = int(count or 0)
tables[table] = row_count
total_rows += row_count
finally:
await engine.dispose()
return {
"source_path": str(path),
"exists": True,
"tables": tables,
"total_rows": total_rows,
}
async def _ensure_target_schema(engine: AsyncEngine) -> None:
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
await conn.run_sync(migrate_accounts_table)
await conn.run_sync(migrate_rules_table)
await conn.run_sync(migrate_message_logs_table)
await conn.run_sync(migrate_users_table)
await conn.run_sync(migrate_payment_orders_table)
await conn.run_sync(migrate_accounts_quota_disabled)
async def _count_target_rows(engine: AsyncEngine) -> dict[str, int]:
counts: dict[str, int] = {}
async with engine.connect() as conn:
for table in MIGRATION_TABLES:
if not await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
counts[table] = 0
continue
count = await conn.scalar(text(f"SELECT COUNT(*) FROM {table}"))
counts[table] = int(count or 0)
return counts
async def _fetch_table_rows(engine: AsyncEngine, table: str) -> list[dict[str, Any]]:
async with engine.connect() as conn:
if not await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
return []
result = await conn.execute(text(f"SELECT * FROM {table} ORDER BY id"))
return [dict(row) for row in result.mappings()]
async def _clear_target_tables(engine: AsyncEngine) -> None:
async with engine.begin() as conn:
dialect = conn.dialect.name
if dialect == "mysql":
await conn.execute(text("SET FOREIGN_KEY_CHECKS=0"))
for table in reversed(MIGRATION_TABLES):
if await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
await conn.execute(text(f"DELETE FROM {table}"))
await conn.execute(text("SET FOREIGN_KEY_CHECKS=1"))
elif dialect == "postgresql":
existing = []
for table in MIGRATION_TABLES:
if await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
existing.append(table)
if existing:
await conn.execute(
text(f"TRUNCATE {', '.join(existing)} RESTART IDENTITY CASCADE")
)
else:
await conn.execute(text("PRAGMA foreign_keys=OFF"))
for table in reversed(MIGRATION_TABLES):
if await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
await conn.execute(text(f"DELETE FROM {table}"))
await conn.execute(text("PRAGMA foreign_keys=ON"))
async def _insert_rows(engine: AsyncEngine, table: str, rows: list[dict[str, Any]]) -> int:
if not rows:
return 0
tbl = Base.metadata.tables.get(table)
if tbl is None:
metadata = MetaData()
def _reflect(sync_conn) -> None:
metadata.reflect(sync_conn, only=[table])
async with engine.connect() as conn:
await conn.run_sync(_reflect)
tbl = metadata.tables[table]
inserted = 0
prepared_rows = _normalize_rows(rows, tbl)
async with engine.begin() as conn:
for offset in range(0, len(prepared_rows), BATCH_SIZE):
batch = prepared_rows[offset : offset + BATCH_SIZE]
await conn.execute(insert(tbl), batch)
inserted += len(batch)
return inserted
async def _reset_auto_increment(engine: AsyncEngine) -> None:
async with engine.begin() as conn:
dialect = conn.dialect.name
for table in MIGRATION_TABLES:
if not await conn.run_sync(lambda sync_conn, t=table: inspect(sync_conn).has_table(t)):
continue
max_id = await conn.scalar(text(f"SELECT COALESCE(MAX(id), 0) FROM {table}"))
if not max_id:
continue
next_id = int(max_id) + 1
if dialect == "mysql":
await conn.execute(text(f"ALTER TABLE {table} AUTO_INCREMENT = {next_id}"))
elif dialect == "postgresql":
await conn.execute(
text(
"SELECT setval("
f"pg_get_serial_sequence('{table}', 'id'), "
f"{next_id}, false)"
)
)
def _build_target_config(payload: dict[str, Any]) -> DatabaseConfig:
current = read_database_config()
password = (payload.get("db_password") or "").strip()
if password in ("", PASSWORD_PLACEHOLDER):
password = current.db_password
return DatabaseConfig(
db_type=payload.get("db_type") or current.db_type,
db_host=(payload.get("db_host") or current.db_host or "127.0.0.1").strip(),
db_port=int(payload.get("db_port") or current.resolved_port() or 0),
db_user=(payload.get("db_user") or current.db_user).strip(),
db_password=password,
db_name=(payload.get("db_name") or current.db_name or "kefu").strip(),
db_path=(payload.get("db_path") or current.db_path).strip(),
)
async def migrate_sqlite_to_target(
payload: dict[str, Any],
*,
clear_target: bool = False,
) -> dict[str, Any]:
source_path = resolve_sqlite_source_path(payload.get("source_db_path"))
if not source_path.exists():
raise FileNotFoundError(f"SQLite 源文件不存在: {source_path}")
target_cfg = _build_target_config(payload)
if _same_sqlite_target(source_path, target_cfg):
raise ValueError("目标库不能与 SQLite 源文件相同,请指定不同的目标数据库")
source_cfg = build_sqlite_config(source_path)
source_engine = create_database_engine(source_cfg)
target_engine = create_database_engine(target_cfg)
try:
preview = await inspect_sqlite_source(str(source_path))
if preview["total_rows"] == 0:
raise ValueError("SQLite 源库中没有可迁移的数据")
await _ensure_target_schema(target_engine)
target_counts = await _count_target_rows(target_engine)
target_total = sum(target_counts.values())
if target_total > 0 and not clear_target:
raise ValueError(
f"目标库已有 {target_total} 条数据,请勾选「清空目标库后再导入」或先手动清空"
)
if clear_target and target_total > 0:
await _clear_target_tables(target_engine)
copied: dict[str, int] = {}
total_rows = 0
for table in MIGRATION_TABLES:
rows = await _fetch_table_rows(source_engine, table)
copied[table] = await _insert_rows(target_engine, table, rows)
total_rows += copied[table]
await _reset_auto_increment(target_engine)
finally:
await source_engine.dispose()
await target_engine.dispose()
target_type = target_cfg.normalized_type()
return {
"message": f"已从 SQLite 迁移 {total_rows} 条记录到 {target_type} 数据库",
"source_path": str(source_path),
"target_type": target_type,
"target_url": mask_database_url(build_database_url(target_cfg)),
"tables": copied,
"total_rows": total_rows,
}