From 6c2f1982615382b317ce176273e095b2a9b72150 Mon Sep 17 00:00:00 2001 From: Simon Date: Tue, 8 Sep 2026 23:58:48 +0800 Subject: [PATCH] =?UTF-8?q?feat(db):=20SQLite=20=E5=85=A8=E9=87=8F?= =?UTF-8?q?=E8=BF=81=E7=A7=BB=E8=87=B3=20MySQL=EF=BC=88config.yaml=20?= =?UTF-8?q?=E9=85=8D=E7=BD=AE=E5=8C=96=20+=20=E8=BF=81=E7=A7=BB=E8=84=9A?= =?UTF-8?q?=E6=9C=AC=20+=20=E4=B8=80=E8=87=B4=E6=80=A7=E6=A0=A1=E9=AA=8C?= =?UTF-8?q?=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 万行迁移并经校验一致 --- .env.example | 10 +- backend/app/core/config.py | 26 ++- backend/pyproject.toml | 1 + backend/tests/conftest.py | 8 +- backend/tests/test_config.py | 77 ++++++-- backend/uv.lock | 11 ++ config.yaml | 16 +- scripts/migrate_sqlite_to_mysql.py | 280 +++++++++++++++++++++++++++++ 8 files changed, 409 insertions(+), 20 deletions(-) create mode 100644 scripts/migrate_sqlite_to_mysql.py diff --git a/.env.example b/.env.example index 701829a..3401faa 100644 --- a/.env.example +++ b/.env.example @@ -9,12 +9,16 @@ TUSHARE_TOKEN= # ---- 数据库连接 ---- -# 留空时使用默认 SQLite:<项目根>/data/quant.db(相对路径自动解析到项目根) +# 默认库由 config.yaml database.mysql 决定(MySQL 192.168.1.10/qlib)。 +# 连接优先级:本文件 DATABASE_URL > config.yaml database.mysql > SQLite 兜底。 DATABASE_URL= # SQLite 示例(显式指定): # DATABASE_URL=sqlite:///./data/quant.db -# MySQL 示例(未来切换,仅需改此值,业务代码不变): -# DATABASE_URL=mysql+pymysql://user:password@127.0.0.1:3306/quant +# MySQL 示例(手动覆盖 config.yaml 的 mysql 段时): +# DATABASE_URL=mysql+pymysql://user:password@127.0.0.1:3306/qlib + +# MySQL 密码(config.yaml database.mysql.password_env 引用;host/port/db/user 在 config.yaml) +MYSQL_PASSWORD= # ---- AI Agent / LLM(Phase 5)---- # 只需填写 API Key;URL 与模型名已在 config.yaml 的 agent.llm 中配置 diff --git a/backend/app/core/config.py b/backend/app/core/config.py index 0c6417a..1921907 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -10,6 +10,7 @@ import os from dataclasses import dataclass from functools import lru_cache from pathlib import Path +from urllib.parse import quote_plus import yaml @@ -99,6 +100,26 @@ def _normalize_sqlite_url(url: str) -> str: return f"{_SQLITE_PREFIX}{(PROJECT_ROOT / rest).resolve()}" +def _build_mysql_url(mysql: dict | None) -> str | None: + """由 config.yaml database.mysql 段组装 mysql+pymysql URL。 + + 规则:仅当 enabled 且 host/user/db 齐全时返回 URL;密码经 password_env + 指定的环境变量读取(AGENT.md §33 密钥只放 .env),未设置则按空密码处理。 + """ + if not mysql or not mysql.get("enabled"): + return None + host = mysql.get("host") + user = mysql.get("user") + db = mysql.get("db") + if not (host and user and db): + return None + port = mysql.get("port") or 3306 + charset = mysql.get("charset") or "utf8mb4" + password = os.environ.get(mysql.get("password_env") or "MYSQL_PASSWORD", "") + auth = f"{quote_plus(user)}:{quote_plus(password)}" if password else quote_plus(user) + return f"mysql+pymysql://{auth}@{host}:{port}/{db}?charset={charset}" + + @lru_cache def get_settings() -> Settings: """加载并缓存 Settings。config / env 路径可通过参数覆盖以便测试。""" @@ -115,7 +136,10 @@ def get_settings() -> Settings: llm_model_env = agent_llm.get("model_env") or "LLM_MODEL" default_url = "sqlite:///./data/quant.db" - database_url = os.environ.get(url_env) or default_url + # URL 优先级:环境变量(DATABASE_URL) > config.yaml database.mysql 段 > SQLite 兜底 + database_url = os.environ.get(url_env) or _build_mysql_url( + _deep(cfg, "database.mysql") or {} + ) or default_url migrations_rel = _deep(cfg, "database.migrations_dir") or ( "app/infrastructure/persistence/migrations" diff --git a/backend/pyproject.toml b/backend/pyproject.toml index b686436..b00619a 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -14,6 +14,7 @@ dependencies = [ # Qlib 研究引擎:PyPI 无 aarch64 wheel(见 docs/ROADMAP.md §2 备注),故从源码 git 安装并固定 commit。 # 网络下载困难时使用 HTTP 代理(见根目录 AGENT.md §0:192.168.1.160:3128)。 "pyqlib @ git+https://github.com/microsoft/qlib.git@79633dd", + "pymysql>=1.1", ] [project.optional-dependencies] diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index fb747f5..af320b7 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -1,8 +1,10 @@ """全局测试配置。 +- 强制 SQLite 临时文件数据库:默认库已是 MySQL(config.yaml database.mysql), + 测试绝不允许触碰开发 MySQL(qlib@192.168.1.10);该 env 在 app 模块 + 首次导入前生效,SessionLocal/engine 会按它构造。 - JOB_MODE=local:单测/CI 里 Job 在本进程执行(不 spawn 子进程、不依赖安装态) - QLIB_SKIP_STALE_JOB_CLEANUP=1:TestClient 启动 lifespan 不清理残留 Job - (避免连接/写入开发库 data/quant.db) 必须在任何 app / config 模块首次导入前生效,故放在 conftest 模块顶层。 """ @@ -11,6 +13,10 @@ from __future__ import annotations import os +# 测试数据库:每进程唯一 /tmp sqlite 文件(隔离、不落项目目录、不触 MySQL) +_TEST_DB = f"/tmp/qlib-pytest-{os.getpid()}.db" +os.environ["DATABASE_URL"] = f"sqlite:///{_TEST_DB}" + # 强制(不是 setdefault):测试绝不 spawn 研究子进程 / 不触碰开发库 os.environ["JOB_MODE"] = "local" os.environ["QLIB_SKIP_STALE_JOB_CLEANUP"] = "1" diff --git a/backend/tests/test_config.py b/backend/tests/test_config.py index 14ad0d1..f56d83a 100644 --- a/backend/tests/test_config.py +++ b/backend/tests/test_config.py @@ -1,9 +1,10 @@ -"""配置层测试:默认 SQLite 相对路径解析、config.yaml / .env 加载约定。""" +"""配置层测试:默认 MySQL URL 组装(config.yaml database.mysql)、.env 加载约定、 +SQLite 兜底逻辑。conftest 已强制 DATABASE_URL=tmp sqlite,本文件内用 +monkeypatch/cache_clear 单独验证「无 env 时走 mysql 段」的分支(只组装不连接)。""" from __future__ import annotations import os -from pathlib import Path from app.core.config import PROJECT_ROOT, _load_env_file, get_settings @@ -14,14 +15,62 @@ def test_project_root_points_to_repo_root() -> None: assert (PROJECT_ROOT / "backend").is_dir() -def test_default_database_url_resolves_to_project_data_dir() -> None: +def test_test_runner_isolation_uses_tmp_sqlite() -> None: + """conftest 强制每进程唯一 /tmp sqlite,测试绝不触碰 MySQL 开发库。""" settings = get_settings() - assert settings.database_url.startswith("sqlite:///") - # 相对路径应解析到 <项目根>/data/quant.db - assert settings.database_url.endswith("/data/quant.db") - db_path = Path(settings.database_url.removeprefix("sqlite:///")) - assert db_path.is_absolute() - assert db_path == PROJECT_ROOT / "data" / "quant.db" + assert settings.database_url.startswith("sqlite:////tmp/qlib-pytest-") + + +def test_mysql_default_from_config_yaml(monkeypatch) -> None: + """未设 DATABASE_URL 时,config.yaml database.mysql 段组装 MySQL URL(不连接)。""" + monkeypatch.delenv("DATABASE_URL", raising=False) + get_settings.cache_clear() + try: + url = get_settings().database_url + finally: + get_settings.cache_clear() + assert url.startswith("mysql+pymysql://qlib:") + assert "192.168.1.10:3306/qlib?charset=utf8mb4" in url + + +def test_build_mysql_url(monkeypatch) -> None: + from app.core.config import _build_mysql_url + + monkeypatch.setenv("MYSQL_PASSWORD", "p@ss:word") + cfg = { + "enabled": True, + "host": "10.0.0.2", + "port": 3307, + "db": "q", + "user": "u", + "password_env": "MYSQL_PASSWORD", + } + url = _build_mysql_url(cfg) + assert url == "mysql+pymysql://u:p%40ss%3Aword@10.0.0.2:3307/q?charset=utf8mb4" + # 未启用 / 缺字段 → None(回退 sqlite 兜底) + assert _build_mysql_url({**cfg, "enabled": False}) is None + assert _build_mysql_url({**cfg, "host": None}) is None + assert _build_mysql_url(None) is None + + +def test_mysql_url_password_optional(monkeypatch) -> None: + """密码留空(未设 env)也可组装 URL —— 方便仅内网/免密场景。""" + from app.core.config import _build_mysql_url + + monkeypatch.delenv("MYSQL_PASSWORD", raising=False) + cfg = {"enabled": True, "host": "h", "db": "d", "user": "u"} + assert _build_mysql_url(cfg) == "mysql+pymysql://u@h:3306/d?charset=utf8mb4" + + +def test_sqlite_url_normalization_and_fallback() -> None: + """仅 sqlite 相对路径被解析为项目根绝对路径;其它 URL 原样透传。""" + from app.core.config import _normalize_sqlite_url + + assert _normalize_sqlite_url("mysql+pymysql://u:p@h/d") == "mysql+pymysql://u:p@h/d" + abs_url = _normalize_sqlite_url("sqlite:///./data/quant.db") + assert abs_url.startswith("sqlite:///") + # 绝对路径 sqlite 不再重复解析 + assert _normalize_sqlite_url("sqlite:////abs/x.db") == "sqlite:////abs/x.db" def test_settings_loaded_from_config_yaml() -> None: @@ -32,7 +81,7 @@ def test_settings_loaded_from_config_yaml() -> None: assert settings.data_source_fallback == "sina" -def test_env_file_loading(tmp_path: Path, monkeypatch) -> None: +def test_env_file_loading(tmp_path, monkeypatch) -> None: monkeypatch.delenv("TUSHARE_TOKEN", raising=False) env_file = tmp_path / ".env" env_file.write_text( @@ -57,14 +106,14 @@ def test_llm_from_config_yaml() -> None: def test_llm_env_overrides_yaml(monkeypatch) -> None: - from app.core.config import get_settings - monkeypatch.setenv("LLM_MODEL", "env-model") monkeypatch.setenv("LLM_BASE_URL", "https://example.com/v1") monkeypatch.setenv("LLM_API_KEY", "sk-test") get_settings.cache_clear() # Settings 有 lru_cache,刷新以读取新环境变量 - s = get_settings() + try: + s = get_settings() + finally: + get_settings.cache_clear() assert s.llm_model == "env-model" assert s.llm_base_url == "https://example.com/v1" assert s.llm_api_key == "sk-test" - get_settings.cache_clear() diff --git a/backend/uv.lock b/backend/uv.lock index 97f1817..88c10b2 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -3262,6 +3262,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/64/02/b2606ddf52fa615d4ff3edb2aa05fb064337fa871f99ebb184a39e051cc8/pymongo-4.18.0-cp314-cp314t-win_arm64.whl", hash = "sha256:4a81d166a43e8af1e5152b6854a263ba0a8831f7dd2ca1badc716f219f4f1bc0", size = 817605, upload-time = "2026-09-03T16:00:41.547Z" }, ] +[[package]] +name = "pymysql" +version = "1.2.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/c9/bc/1c6a92f385940f727daeecf3bacaf186e03875dff57197801046c583bcf0/pymysql-1.2.0.tar.gz", hash = "sha256:6c7b17ca686988104d7426c27895b455cdeea3e9d3ceb1270f0c3704fead8c33", size = 49021, upload-time = "2026-05-19T08:26:22.302Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c4/bd/2534e130295c8cfd4f0a2e31623baab7502278f1e97bcfe61db75656a77f/pymysql-1.2.0-py3-none-any.whl", hash = "sha256:62169ce6d5510f08e140c5e7990ee884a9764024e4a9a27b2cc11f1099322ae0", size = 45716, upload-time = "2026-05-19T08:26:20.974Z" }, +] + [[package]] name = "pyparsing" version = "3.3.2" @@ -3540,6 +3549,7 @@ dependencies = [ { name = "alembic" }, { name = "fastapi" }, { name = "httpx" }, + { name = "pymysql" }, { name = "pyqlib" }, { name = "pyyaml" }, { name = "sqlalchemy" }, @@ -3563,6 +3573,7 @@ requires-dist = [ { name = "alembic", specifier = ">=1.13" }, { name = "fastapi", specifier = ">=0.115" }, { name = "httpx", specifier = ">=0.28.1" }, + { name = "pymysql", specifier = ">=1.1" }, { name = "pyqlib", git = "https://github.com/microsoft/qlib.git?rev=79633dd" }, { name = "pyyaml", specifier = ">=6.0" }, { name = "sqlalchemy", specifier = ">=2.0" }, diff --git a/config.yaml b/config.yaml index 283089e..b4321fb 100644 --- a/config.yaml +++ b/config.yaml @@ -17,11 +17,25 @@ api: prefix: "/api" database: - # 留空 → 默认 sqlite:///./data/quant.db(相对路径自动解析到项目根 data/) + # 数据库 URL 选择优先级: + # 1) 环境变量 database.url_env(如 DATABASE_URL,可在 .env 覆盖) + # 2) database.mysql(enabled: true 时由 host/port/db/user 组装 mysql+pymysql URL) + # 3) 兜底 sqlite:///./data/quant.db(相对路径自动解析到项目根 data/) url_env: "DATABASE_URL" echo: false # Alembic 迁移脚本目录(相对 backend/) migrations_dir: "app/infrastructure/persistence/migrations" + # ---- MySQL(当前默认库;2026-09 由 SQLite data/quant.db 全量迁移而来)---- + # host/port/db/user 明文可提交;密码经 password_env 引用 .env 变量 + # (AGENT.md §33:密钥只放根目录 .env,本文件不写明文密码)。 + mysql: + enabled: true + host: "192.168.1.10" + port: 3306 + db: "qlib" + user: "qlib" + password_env: "MYSQL_PASSWORD" + charset: "utf8mb4" data_source: primary: "tushare" diff --git a/scripts/migrate_sqlite_to_mysql.py b/scripts/migrate_sqlite_to_mysql.py new file mode 100644 index 0000000..f6d5aeb --- /dev/null +++ b/scripts/migrate_sqlite_to_mysql.py @@ -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())