Files
qlib/scripts/migrate_sqlite_to_mysql.py
Simon 6c2f198261 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 万行迁移并经校验一致
2026-09-08 23:58:48 +08:00

281 lines
11 KiB
Python
Raw Permalink 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.
#!/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())