feat(db): SQLite 全量迁移至 MySQL(config.yaml 配置化 + 迁移脚本 + 一致性校验)
- config.yaml database.mysql:host/port/db/user/charset 明文可提交;密码经 password_env 引用 .env 的 MYSQL_PASSWORD(AGENT.md §33 密钥不进 git) - config.py _build_mysql_url:URL 优先级 DATABASE_URL env > database.mysql 段 > sqlite 兜底 - pyproject 引入 pymysql>=1.1 - tests/conftest.py 强制每进程 /tmp SQLite(测试绝不触 MySQL 开发库);test_config 覆盖 mysql 组装/密码可选/sqlite 兜底分支 - scripts/migrate_sqlite_to_mysql.py:sqlite→mysql 一次性迁移工具(keyset 分页 + chunk 多值 INSERT + 攒批 commit + 幂等续传 + 低配 MySQL 节流 --throttle-sec + --verify-only 行数与抽样一致性校验);已用于 data/quant.db 约 1605 万行迁移并经校验一致
This commit is contained in:
@@ -0,0 +1,280 @@
|
||||
#!/usr/bin/env python3
|
||||
"""SQLite → MySQL 全量数据迁移脚本(一次性工具,不入应用代码)。
|
||||
|
||||
把 <项目根>/data/quant.db(SQLite)的全部业务表数据迁移到 MySQL 目标库,
|
||||
供后端以 MySQL 作为默认数据库运行(docs/DEV_PLAN「数据库迁移」专项)。
|
||||
|
||||
设计:
|
||||
- 源:标准库 sqlite3 直连(只读查询)。
|
||||
- 目标:SQLAlchemy engine(URL 形态任意,mysql+pymysql 由配置/env 提供),
|
||||
用 raw_connection() 拿到 pymysql 连接做 chunk 多值 INSERT 批量写入。
|
||||
- 逐表 keyset 分页(int 主键按 id > last),保留原始主键 id。
|
||||
- 幂等:目标表已有数据时从 max(id)+1 续传;string 主键小表整表一次。
|
||||
选项 --reset 先 TRUNCATE 再全量重传。
|
||||
- 完成后逐表 COUNT 校验 + 抽样逐行比对。
|
||||
|
||||
用法:
|
||||
cd backend && .venv/bin/python ../scripts/migrate_sqlite_to_mysql.py [--reset] [--verify-only]
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sqlite3
|
||||
import sys
|
||||
import time
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
BACKEND = ROOT / "backend"
|
||||
DEFAULT_SOURCE = ROOT / "data" / "quant.db"
|
||||
|
||||
# 表 → (列清单, 主键类型)。列顺序即插入顺序;全部显式带 id。
|
||||
TABLES: dict[str, tuple[list[str], str]] = {
|
||||
"stock": (
|
||||
["id", "symbol", "name", "industry", "area", "market", "exchange",
|
||||
"list_date", "delist_date", "status"],
|
||||
"int",
|
||||
),
|
||||
"trading_calendar": (["id", "calendar_date", "is_open"], "int"),
|
||||
"sync_log": (
|
||||
["id", "source", "api", "request_time", "success", "failure_reason",
|
||||
"row_count", "data_start", "data_end"],
|
||||
"int",
|
||||
),
|
||||
"job": (
|
||||
["id", "kind", "status", "stage", "spec_json", "error", "result_json",
|
||||
"experiment_id", "created_at", "started_at", "finished_at"],
|
||||
"str",
|
||||
),
|
||||
"experiment": (
|
||||
["id", "kind", "spec_json", "result_json", "summary_text",
|
||||
"code_version", "data_version", "job_id", "created_at"],
|
||||
"str",
|
||||
),
|
||||
"financial_indicator": (
|
||||
["id", "symbol", "report_date", "announce_date", "source", "eps", "roe",
|
||||
"total_revenue", "net_profit", "gross_margin"],
|
||||
"int",
|
||||
),
|
||||
"stock_daily": (
|
||||
["id", "symbol", "trade_date", "source", "adjust", "open", "high",
|
||||
"low", "close", "volume", "amount"],
|
||||
"int",
|
||||
),
|
||||
"adjust_factor": (["id", "symbol", "trade_date", "factor"], "int"),
|
||||
}
|
||||
|
||||
CHUNK = 5000 # 每批行数:5000 × 11 列占位符 ≈ 5.5 万 < MySQL 65535 上限
|
||||
PROGRESS_EVERY = 200_000
|
||||
|
||||
|
||||
def resolve_dst_engine():
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
url = os.environ.get("DATABASE_URL")
|
||||
if not url:
|
||||
sys.path.insert(0, str(BACKEND))
|
||||
from app.core.config import get_settings
|
||||
|
||||
url = get_settings().database_url
|
||||
print(f"目标 URL: {url.split('@')[-1]}", flush=True)
|
||||
return create_engine(url, future=True, pool_pre_ping=True)
|
||||
|
||||
|
||||
def migrate_table(src: sqlite3.Connection, dst, table: str, cols: list[str], pk_kind: str,
|
||||
*, reset: bool, throttle_sec: float = 0.0) -> int:
|
||||
col_list = ", ".join(cols) # sqlite 侧
|
||||
col_q = ", ".join(f"`{c}`" for c in cols) # mysql 侧
|
||||
src_total = src.execute(f'SELECT COUNT(*) FROM "{table}"').fetchone()[0]
|
||||
|
||||
cur = dst.cursor()
|
||||
try:
|
||||
if pk_kind == "str":
|
||||
# 小表(job/experiment):整表读入,幂等靠唯一主键冲突忽略
|
||||
if reset:
|
||||
cur.execute(f"TRUNCATE TABLE `{table}`")
|
||||
dst.commit()
|
||||
elif count(dst, table):
|
||||
print(f"[{table}] 目标非空且未 --reset,跳过(校验阶段会核对)")
|
||||
return 0
|
||||
rows = src.execute(f'SELECT {col_list} FROM "{table}"').fetchall()
|
||||
if rows:
|
||||
_insert_chunk(cur, table, col_q, cols, rows)
|
||||
dst.commit()
|
||||
print(f"[{table}] 完成 {len(rows)} 行")
|
||||
return len(rows)
|
||||
|
||||
# int 主键:keyset 分页
|
||||
if reset:
|
||||
cur.execute(f"TRUNCATE TABLE `{table}`")
|
||||
dst.commit()
|
||||
last = 0
|
||||
else:
|
||||
cur.execute(f"SELECT COALESCE(MAX(id), 0) FROM `{table}`")
|
||||
last = int(cur.fetchone()[0])
|
||||
if last:
|
||||
print(f"[{table}] 目标已有数据,从 id={last + 1} 续传")
|
||||
inserted = 0
|
||||
t0 = time.time()
|
||||
commit_every = max(100_000 // CHUNK, 1) * CHUNK # 每 ~10 万行 commit 一次
|
||||
while True:
|
||||
rows = src.execute(
|
||||
f'SELECT {col_list} FROM "{table}" WHERE id > ? ORDER BY id LIMIT ?',
|
||||
(last, CHUNK),
|
||||
).fetchall()
|
||||
if not rows:
|
||||
break
|
||||
_insert_chunk(cur, table, col_q, cols, rows) # execute 但不立即 commit
|
||||
inserted += len(rows)
|
||||
last = rows[-1][0]
|
||||
if throttle_sec:
|
||||
time.sleep(throttle_sec) # 低配 MySQL:每批间限速,降低写入压力
|
||||
if inserted % commit_every < CHUNK:
|
||||
dst.commit() # 攒批 commit:显著减少 fsync 次数
|
||||
if inserted // PROGRESS_EVERY != (inserted - len(rows)) // PROGRESS_EVERY:
|
||||
el = time.time() - t0
|
||||
rate = inserted / el if el else 0
|
||||
print(f"[{table}] {inserted:,}/{src_total:,} 行 ({el:.0f}s, {rate:,.0f} 行/s)",
|
||||
flush=True)
|
||||
dst.commit()
|
||||
el = time.time() - t0
|
||||
print(f"[{table}] 完成 {inserted:,} 行,耗时 {el:.0f}s", flush=True)
|
||||
return inserted
|
||||
finally:
|
||||
cur.close()
|
||||
|
||||
|
||||
def _insert_chunk(cur, table: str, col_q: str, cols: list[str], rows) -> None:
|
||||
n = len(rows)
|
||||
n_col = len(cols)
|
||||
per = ", ".join(["%s"] * n_col)
|
||||
values = ", ".join([f"({per})"] * n)
|
||||
sql = f"INSERT INTO `{table}` ({col_q}) VALUES {values}"
|
||||
flat = [v for row in rows for v in row]
|
||||
cur.execute(sql, flat)
|
||||
|
||||
|
||||
def count(dst, table: str) -> int:
|
||||
cur = dst.cursor()
|
||||
try:
|
||||
cur.execute(f"SELECT COUNT(*) FROM `{table}`")
|
||||
return int(cur.fetchone()[0])
|
||||
finally:
|
||||
cur.close()
|
||||
|
||||
|
||||
def _verify_eq(a, b) -> bool:
|
||||
"""抽样比对:None 严格相等;日期对象与 ISO 字符串互比;数值近似
|
||||
(sqlite float vs mysql Decimal 舍入差);其余直等。"""
|
||||
if a is None or b is None:
|
||||
return a is b
|
||||
# 数值(sqlite float / mysql Decimal / int):按 6 位小数容差比较
|
||||
if isinstance(a, (int, float, Decimal)) and isinstance(b, (int, float, Decimal)):
|
||||
return round(float(a), 6) == round(float(b), 6)
|
||||
if isinstance(a, str) and hasattr(b, "isoformat"):
|
||||
return a == b.isoformat()
|
||||
if isinstance(b, str) and hasattr(a, "isoformat"):
|
||||
return a.isoformat() == b
|
||||
return a == b
|
||||
|
||||
|
||||
def verify(src: sqlite3.Connection, dst) -> bool:
|
||||
ok = True
|
||||
print("---- 行数校验 ----", flush=True)
|
||||
for table, (_, _pk) in TABLES.items():
|
||||
s = src.execute(f'SELECT COUNT(*) FROM "{table}"').fetchone()[0]
|
||||
d = count(dst, table)
|
||||
flag = "OK" if s == d else "MISMATCH"
|
||||
ok = ok and s == d
|
||||
print(f" {table:<22} sqlite={s:>12,} mysql={d:>12,} {flag}", flush=True)
|
||||
|
||||
print("---- 抽样比对 ----", flush=True)
|
||||
cur = dst.cursor()
|
||||
try:
|
||||
cases = [
|
||||
# 注意:sqlite/mysql 双兼容写法 —— 表名裸名 + 字符串一律单引号
|
||||
# (mysql 版仅把表名包反引号,见下方 replace("FROM ", "FROM `") 思路外实现)
|
||||
("stock_daily",
|
||||
"SELECT symbol, trade_date, open, high, low, close, volume, amount, source, adjust "
|
||||
'FROM "stock_daily" WHERE symbol IN (\'600519.SH\',\'000001.SZ\') '
|
||||
"ORDER BY trade_date, symbol LIMIT 300"),
|
||||
("adjust_factor",
|
||||
'SELECT symbol, trade_date, factor FROM "adjust_factor" WHERE symbol = \'600519.SH\' '
|
||||
"ORDER BY trade_date, symbol LIMIT 300"),
|
||||
("financial_indicator",
|
||||
'SELECT symbol, report_date, announce_date, eps, roe, total_revenue, net_profit '
|
||||
'FROM "financial_indicator" WHERE symbol = \'600519.SH\' '
|
||||
"ORDER BY report_date, symbol LIMIT 300"),
|
||||
("stock",
|
||||
'SELECT symbol, name, industry, market, list_date, status FROM "stock" '
|
||||
"WHERE symbol IN ('600519.SH','000001.SZ','300750.SZ')"),
|
||||
]
|
||||
for table, ssql in cases:
|
||||
s_rows = src.execute(ssql).fetchall()
|
||||
# mysql 版:双引号仅用于表名(字符串已是单引号),安全替换为反引号
|
||||
m_sql = ssql.replace('"', "`")
|
||||
cur.execute(m_sql)
|
||||
m_rows = cur.fetchall()
|
||||
bad = 0
|
||||
for a, b in zip(s_rows, m_rows, strict=False):
|
||||
if any(not _verify_eq(x, y) for x, y in zip(a, b, strict=False)):
|
||||
bad += 1
|
||||
same_len = len(s_rows) == len(m_rows)
|
||||
flag = "OK" if bad == 0 and same_len else "MISMATCH"
|
||||
ok = ok and bad == 0 and same_len
|
||||
print(f" {table:<22} sqlite={len(s_rows):>4} mysql={len(m_rows):>4} 不一致行={bad} {flag}",
|
||||
flush=True)
|
||||
finally:
|
||||
cur.close()
|
||||
return ok
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser(description="SQLite → MySQL 数据迁移")
|
||||
ap.add_argument("--source", default=str(DEFAULT_SOURCE))
|
||||
ap.add_argument("--reset", action="store_true", help="目标表已有数据时 TRUNCATE 重传")
|
||||
ap.add_argument("--verify-only", action="store_true", help="仅校验不写入")
|
||||
ap.add_argument("--only", nargs="+", default=None,
|
||||
help="只迁移指定表(可多个);并行迁移大表时使用(自动跳过末尾校验)")
|
||||
ap.add_argument("--throttle-sec", type=float, default=0.0,
|
||||
help="每批(5000 行)之间的休眠秒数;低配 MySQL/共享服务器请设 0.3~1.0 降压力")
|
||||
args = ap.parse_args()
|
||||
|
||||
tables = {k: v for k, v in TABLES.items() if args.only is None or k in args.only}
|
||||
if not tables:
|
||||
print(f"--only 未匹配任何表,可用: {', '.join(TABLES)}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
src_path = Path(args.source)
|
||||
if not src_path.exists():
|
||||
print(f"源 sqlite 不存在: {src_path}", file=sys.stderr)
|
||||
return 2
|
||||
print(f"源 : {src_path} ({src_path.stat().st_size / 1e9:.2f} GB)", flush=True)
|
||||
print(f"目标库: {resolve_dst_engine().url.database}", flush=True)
|
||||
|
||||
src = sqlite3.connect(str(src_path))
|
||||
engine = resolve_dst_engine()
|
||||
dst = engine.raw_connection()
|
||||
try:
|
||||
if args.verify_only:
|
||||
ok = verify(src, dst)
|
||||
return 0 if ok else 1
|
||||
for table, (cols, pk_kind) in tables.items():
|
||||
migrate_table(src, dst, table, cols, pk_kind, reset=args.reset,
|
||||
throttle_sec=args.throttle_sec)
|
||||
if args.only:
|
||||
return 0 # 并行分路:由单独 --verify-only 统一校验
|
||||
ok = verify(src, dst)
|
||||
print("迁移完成:", "全部一致 ✓" if ok else "存在不一致,请人工核对 ✗")
|
||||
return 0 if ok else 1
|
||||
finally:
|
||||
src.close()
|
||||
dst.close()
|
||||
engine.dispose()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user