diff --git a/CHANGELOG.md b/CHANGELOG.md index 16f8d0e..d131084 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,23 @@ 本文件记录 easy-tdx 的版本变更。格式遵循 [Keep a Changelog](https://keepachangelog.com/zh-CN/)。 +## [1.26.0] — 2026-09-01 + +**本地数据仓库版本**——把碎片化缓存升级为统一数据底座(升级计划 P2 阶段;P2-2 评级后端化已随 1.25.0 提前交付)。此前下游项目(indicator-lab 的 DuckDB 仓库、backtest-system 的 cache/ 目录)都在自建数据层,现在 easy-tdx 原生提供。 + +### 新增 + +- **K 线仓库**(`warehouse/` 包,DuckDB 单文件)——默认 `~/.easy_tdx/warehouse.duckdb`(随 `EASY_TDX_CONFIG_DIR`),列存 + SQL 友好 + 主键去重 upsert。DuckDB 为**可选依赖**(`pip install easy-tdx[warehouse]`),惰性导入不影响核心三通道。 +- **provisional / completed 状态机**(借鉴 indicator-lab)——15:05 前落盘的当日 bar **逐行**标记 `provisional`(盘中临时值),查询/回测默认忽略(杜绝拿盘中价当收盘价);`promote_provisional()` 把过期临时行转正,`include_provisional=True` 显式可见。 +- **增量同步器**(`warehouse/sync.py`)——首同步全量(默认上限 8000 根),此后只拉尾部 15 根覆盖(收盘价修正/临时转正),不动更早历史;批次同步带进度回调、单标失败不中断批次,返回 added/updated/skipped/failed 汇总。默认 QFQ 口径(回测/筛选一致)。 +- **仓库健康自检**(`health_check`)——三维度体检:①疑似缺口(相邻 bar 工作日差 > 5,含节假日误报提示);②异常跳变(复用 QFQ 对拍的板块感知跳空检测,多为除权数据需人工核查);③最新度(>7 天未更新的过期标的)+ provisional 行统计。 +- **CLI 命令组** `easy-tdx warehouse`——`sync`(支持逗号分隔或 @文件 标的列表)、`query`(JSON 输出,`--include-provisional`)、`stats`(各标的行数/范围/临时行)、`check`(健康自检)。 + +### 内部 + +- `pyproject.toml` 新增 `[warehouse]` 可选依赖组;`duckdb` 加入 dev 依赖(CI 跑仓库测试)。 +- `cli/__init__.py` 注册 `warehouse` 命令组(37+1 个顶级命令)。 + ## [1.25.0] — 2026-09-01 **防过拟合验证链版本**——补上两个下游项目(backtest-system / indicator-lab)都在自研的最大空白:样本外验证工具链。此后「回测好」可升级为「样本外也好」。升级计划第二阶段(P1),全量 1193 单测。 diff --git a/pyproject.toml b/pyproject.toml index 3c32557..d2d754d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "easy-tdx" -version = "1.25.0" +version = "1.26.0" description = "通达信 TCP 协议行情数据客户端,支持在线行情、离线数据读取与写入同步" readme = "README.md" requires-python = ">=3.10" @@ -14,9 +14,11 @@ dependencies = ["pandas>=2.0,<3", "tzdata>=2024.1", "click>=8.0,<9"] easy-tdx = "easy_tdx.cli:cli" # cli/__init__.py exposes the click group [project.optional-dependencies] -dev = ["pytest>=8.0", "pytest-asyncio>=0.23", "pytest-cov", "mypy>=1.9", "ruff>=0.4", "scipy>=1.10,<1.16", "httpx>=0.27"] +dev = ["pytest>=8.0", "pytest-asyncio>=0.23", "pytest-cov", "mypy>=1.9", "ruff>=0.4", "scipy>=1.10,<1.16", "httpx>=0.27", "duckdb>=1.0"] science = ["scipy>=1.10,<1.16"] web = ["fastapi>=0.110,<1", "uvicorn[standard]>=0.29"] +# 本地 K 线数据仓库(easy-tdx warehouse ...):DuckDB 单文件列存 +warehouse = ["duckdb>=1.0"] # 打包成桌面 EXE 用:系统托盘(pystray)+ 图标生成(Pillow)。 # 仅 PyInstaller 打包态需要,开发态 ``pip install -e .[web]`` 不强制装。 packaging = ["pystray>=0.19", "Pillow>=10.0"] diff --git a/src/easy_tdx/cli/__init__.py b/src/easy_tdx/cli/__init__.py index a2f3e6c..f7bd32e 100644 --- a/src/easy_tdx/cli/__init__.py +++ b/src/easy_tdx/cli/__init__.py @@ -33,6 +33,7 @@ from .cmd_quote import quote, quote_list from .cmd_run_all import run_all from .cmd_tick import tick from .cmd_transaction import transaction +from .cmd_warehouse import warehouse from .cmd_web import serve @@ -96,3 +97,4 @@ cli.add_command(portfolio) cli.add_command(run_all) cli.add_command(screen) cli.add_command(serve) +cli.add_command(warehouse) diff --git a/src/easy_tdx/cli/cmd_warehouse.py b/src/easy_tdx/cli/cmd_warehouse.py new file mode 100644 index 0000000..654f87e --- /dev/null +++ b/src/easy_tdx/cli/cmd_warehouse.py @@ -0,0 +1,172 @@ +"""``easy-tdx warehouse`` 命令组:本地 K 线仓库的同步 / 查询 / 统计 / 自检。 + +仓库为 DuckDB 单文件(默认 ``~/.easy_tdx/warehouse.duckdb``),依赖可选 +安装:``pip install easy-tdx[warehouse]``。 + +示例:: + + easy-tdx warehouse sync --symbols SH:600519,SZ:000001 + easy-tdx warehouse query SH 600519 --count 30 + easy-tdx warehouse stats + easy-tdx warehouse check --symbols SH:600519 +""" + +from __future__ import annotations + +import json +import sys +from typing import Any + +import click + + +def _require_warehouse(db_path: str | None) -> Any: + """惰性导入 warehouse(duckdb 可选依赖),失败给友好错误。""" + try: + from easy_tdx.warehouse import KlineWarehouse + except ImportError as exc: + click.echo(f"错误: {exc}", err=True) + raise SystemExit(2) from exc + return KlineWarehouse(db_path) if db_path else KlineWarehouse() + + +@click.group("warehouse") +def warehouse() -> None: + """本地 K 线数据仓库(DuckDB):同步、查询、统计、健康自检。""" + + +@warehouse.command("sync") +@click.option( + "--symbols", required=True, help="标的列表:逗号分隔(SH:600519,SZ:000001)或 @文件(每行一个)" +) +@click.option("--period", default="DAILY", help="K 线周期(默认 DAILY)") +@click.option("--max-bars", default=8000, type=int, help="首同步最大拉取根数(默认 8000)") +@click.option("--tail-bars", default=15, type=int, help="增量同步尾部根数(默认 15)") +@click.option( + "--adjust", + default="QFQ", + type=click.Choice(["NONE", "QFQ", "HFQ"]), + help="复权口径(默认 QFQ)", +) +@click.option( + "--db", "db_path", default=None, help="仓库文件路径(默认 ~/.easy_tdx/warehouse.duckdb)" +) +def warehouse_sync( + symbols: str, + period: str, + max_bars: int, + tail_bars: int, + adjust: str, + db_path: str | None, +) -> None: + """增量同步行情进仓库(首同步全量、此后只补尾部)。""" + if symbols.startswith("@"): + from pathlib import Path + + path = Path(symbols[1:]) + if not path.exists(): + click.echo(f"错误: 标的列表文件不存在 {path}", err=True) + raise SystemExit(1) + symbol_list = [ + ln.strip() + for ln in path.read_text(encoding="utf-8").splitlines() + if ln.strip() and not ln.strip().startswith("#") + ] + else: + symbol_list = [s.strip() for s in symbols.split(",") if s.strip()] + + with _require_warehouse(db_path) as wh: + from easy_tdx.warehouse import WarehouseSyncer + + def _progress(done: int, total: int, sym: str) -> None: + click.echo(f"[{done}/{total}] {sym}", err=True) + + from ..cli.conn import get_mac_client + + with get_mac_client() as client: + syncer = WarehouseSyncer( + client, wh, max_bars=max_bars, tail_bars=tail_bars, adjust=adjust + ) + summary = syncer.sync(symbol_list, period=period, progress=_progress) + click.echo( + json.dumps( + {k: v for k, v in summary.items() if k != "details"}, + ensure_ascii=False, + ) + ) + + +@warehouse.command("query") +@click.argument("market", type=click.Choice(["SH", "SZ", "BJ"])) +@click.argument("code") +@click.option("--period", default="DAILY", help="K 线周期(默认 DAILY)") +@click.option("--count", default=None, type=int, help="最近 N 根(默认全部)") +@click.option("--start", default=None, help="开始日期 YYYY-MM-DD") +@click.option("--end", default=None, help="结束日期 YYYY-MM-DD") +@click.option("--include-provisional", is_flag=True, help="包含未收盘的临时 bar(默认忽略)") +@click.option("--db", "db_path", default=None, help="仓库文件路径") +def warehouse_query( + market: str, + code: str, + period: str, + count: int | None, + start: str | None, + end: str | None, + include_provisional: bool, + db_path: str | None, +) -> None: + """查询仓库中的 K 线(JSON 输出,datetime/OHLCV/status)。""" + with _require_warehouse(db_path) as wh: + df = wh.query( + market, + code, + period=period, + start=start, + end=end, + count=count, + include_provisional=include_provisional, + ) + if len(df) == 0: + click.echo("[]") + return + out = json.loads(df.to_json(orient="records", date_format="iso", force_ascii=False)) + click.echo(json.dumps(out, ensure_ascii=False)) + + +@warehouse.command("stats") +@click.option("--period", default="DAILY", help="K 线周期(默认 DAILY)") +@click.option("--db", "db_path", default=None, help="仓库文件路径") +def warehouse_stats(period: str, db_path: str | None) -> None: + """仓库统计:各标的行数 / 数据范围 / provisional 行数。""" + with _require_warehouse(db_path) as wh: + df = wh.symbols(period=period) + if len(df) == 0: + click.echo("仓库为空,请先执行 easy-tdx warehouse sync", err=True) + raise SystemExit(1) + df_out = df.copy() + df_out["first"] = df_out["first"].astype(str) + df_out["last"] = df_out["last"].astype(str) + click.echo( + json.dumps( + json.loads(df_out.to_json(orient="records", force_ascii=False)), ensure_ascii=False + ) + ) + + +@warehouse.command("check") +@click.option("--symbols", default=None, help="只检查指定标的(逗号分隔),默认全仓库") +@click.option("--db", "db_path", default=None, help="仓库文件路径") +def warehouse_check(symbols: str | None, db_path: str | None) -> None: + """仓库健康自检:缺口 / 异常跳变 / 最新度 / 临时行。""" + with _require_warehouse(db_path) as wh: + market = code = None + if symbols: + first = [s.strip() for s in symbols.split(",") if s.strip()][0] + market, code = first.split(":", 1) + if ":" not in symbols and len(symbols.split(",")) > 1: + click.echo("错误: --symbols 自检模式一次只支持一个标的", err=True) + raise SystemExit(1) + report = wh.health_check(market=market, code=code) + click.echo(json.dumps(report, ensure_ascii=False, default=str)) + if report["issues"]: + sys.exit(0) # 有问题不报错退出——自检结果本身是正常输出 diff --git a/src/easy_tdx/warehouse/__init__.py b/src/easy_tdx/warehouse/__init__.py new file mode 100644 index 0000000..a8f5d1d --- /dev/null +++ b/src/easy_tdx/warehouse/__init__.py @@ -0,0 +1,21 @@ +"""easy_tdx.warehouse — K 线本地数据仓库(DuckDB,可选依赖)。 + +快速开始:: + + from easy_tdx.warehouse import KlineWarehouse, WarehouseSyncer + + wh = KlineWarehouse() # ~/.easy_tdx/warehouse.duckdb + with get_mac_client() as client: # 首次:全量;此后:增量补缺 + WarehouseSyncer(client, wh).sync(["SH:600519", "SZ:000001"]) + df = wh.query("SH", "600519", count=250) # 默认忽略未收盘 provisional bar +""" + +from easy_tdx.warehouse.store import KlineWarehouse, default_warehouse_path +from easy_tdx.warehouse.sync import SyncSummary, WarehouseSyncer + +__all__ = [ + "KlineWarehouse", + "WarehouseSyncer", + "SyncSummary", + "default_warehouse_path", +] diff --git a/src/easy_tdx/warehouse/store.py b/src/easy_tdx/warehouse/store.py new file mode 100644 index 0000000..fa137d0 --- /dev/null +++ b/src/easy_tdx/warehouse/store.py @@ -0,0 +1,423 @@ +"""K 线本地数据仓库(DuckDB,v1.26 新增)。 + +此前 easy-tdx 的缓存是碎片化的(股票列表缓存 / best_host / 进程内 XDXR +字典 / 扫描 JSON),没有统一的 K 线磁盘层——下游项目(indicator-lab 的 +DuckDB 仓库、backtest-system 的 cache/ 目录)都在自建数据层。本模块把 +「拉过的行情」沉淀为本地列存仓库: + +- **单文件 DuckDB**(默认 ``~/.easy_tdx/warehouse.duckdb``,随 + ``EASY_TDX_CONFIG_DIR`` 走):零服务、列存、SQL 友好、压缩率高 + (万只标的十年日线约几百 MB); +- **增量同步**:每标的只补最后若干根(首同步全量拉取),同日 bar 用新值 + 覆盖(收盘修正),历史数据不动; +- **provisional / completed 状态机**(借鉴 indicator-lab):15:05 前落盘 + 的当日 bar 标记 ``provisional``(未收盘的临时值),查询/回测**默认忽略**, + 只在显式 ``include_provisional=True`` 时可见;次日增量同步自动把过期 + 临时行转正/覆盖。 + +DuckDB 为可选依赖(``pip install easy-tdx[warehouse]``),惰性导入—— +未安装时本模块给出明确报错,不影响核心三通道。 +""" + +from __future__ import annotations + +import logging +from datetime import date, datetime +from datetime import time as dt_time +from pathlib import Path +from typing import Any + +import pandas as pd + +logger = logging.getLogger(__name__) + +__all__ = ["KlineWarehouse", "default_warehouse_path", "MARKET_TO_TDX"] + +# 未收盘 cutoff:15:05(A股 15:00 收盘 + 5 分钟数据落定余量) +_MARKET_CLOSE_CUTOFF = dt_time(15, 5) + +MARKET_TO_TDX: dict[str, int] = {"SZ": 0, "SH": 1, "BJ": 2} +_TDX_TO_MARKET: dict[int, str] = {v: k for k, v in MARKET_TO_TDX.items()} + +_SCHEMA = """ +CREATE TABLE IF NOT EXISTS klines ( + market TEXT NOT NULL, + code TEXT NOT NULL, + period TEXT NOT NULL, + datetime TIMESTAMP NOT NULL, + open DOUBLE, + high DOUBLE, + low DOUBLE, + close DOUBLE, + vol DOUBLE, + amount DOUBLE, + status TEXT NOT NULL DEFAULT 'completed', + updated_at TIMESTAMP, + PRIMARY KEY (market, code, period, datetime) +); +CREATE INDEX IF NOT EXISTS idx_klines_symbol ON klines (market, code, period, datetime); +""" + + +def default_warehouse_path() -> Path: + """默认仓库文件路径(随 ``EASY_TDX_CONFIG_DIR``)。""" + import os + + base = Path(os.environ.get("EASY_TDX_CONFIG_DIR", str(Path.home() / ".easy_tdx"))) + return base / "warehouse.duckdb" + + +def _require_duckdb() -> Any: + try: + import duckdb # noqa: PLC0415 — 惰性导入,可选依赖 + except ImportError as exc: + raise ImportError( + "warehouse 功能需要 DuckDB:pip install easy-tdx[warehouse](或 pip install duckdb)" + ) from exc + return duckdb + + +class KlineWarehouse: + """K 线 DuckDB 仓库:upsert / 查询 / provisional 状态机 / 健康自检。 + + Example:: + wh = KlineWarehouse() # 或 KlineWarehouse(path) + wh.upsert_bars("SH", "600519", df) # 增量写入 + df = wh.query("SH", "600519", count=250) # 默认忽略 provisional + """ + + def __init__(self, db_path: Path | str | None = None) -> None: + self._duckdb = _require_duckdb() + self._path = Path(db_path) if db_path is not None else default_warehouse_path() + self._path.parent.mkdir(parents=True, exist_ok=True) + self._conn = self._duckdb.connect(str(self._path)) + self._conn.execute(_SCHEMA) + + # ── 基本属性 ───────────────────────────────────────────────────────────── + + @property + def path(self) -> Path: + return self._path + + def close(self) -> None: + """关闭连接。""" + try: + self._conn.close() + except Exception: # noqa: BLE001 — 幂等关闭 + pass + + def __enter__(self) -> KlineWarehouse: + return self + + def __exit__(self, *args: Any) -> None: + self.close() + + # ── 写入 ───────────────────────────────────────────────────────────────── + + def upsert_bars( + self, + market: str, + code: str, + df: pd.DataFrame, + period: str = "DAILY", + status: str | None = None, + ) -> tuple[int, int]: + """Upsert 一批 K 线,返回 ``(新增数, 更新数)``。 + + Args: + market: ``"SH"`` / ``"SZ"`` / ``"BJ"``。 + code: 6 位代码。 + df: 含 ``datetime``(或 ``date``)与 OHLCV 列的 K 线。 + period: 周期名(默认 ``"DAILY"``,与 Period 名对齐)。 + status: 显式状态;None = 自动(当日 bar 且未过 15:05 → + ``provisional``,其余 ``completed``)。 + """ + if df is None or len(df) == 0: + return (0, 0) + mkt = market.upper() + src = df.copy() + dt_col = "datetime" if "datetime" in src.columns else "date" + src["_dt"] = pd.to_datetime(src[dt_col]) + cols = ["open", "high", "low", "close", "vol", "amount"] + for c in cols: + if c not in src.columns: + src[c] = float("nan") + + now = datetime.now() + today = now.date() + before_close = now.time() < _MARKET_CLOSE_CUTOFF + + def _row_status(ts: pd.Timestamp) -> str: + """逐行判定:当日 bar 且未过 15:05 → provisional,否则 completed。""" + if status is not None: + return status + return "provisional" if (ts.date() == today and before_close) else "completed" + + # 统计新增/更新(按已有键) + existing = self._conn.execute( + """ + SELECT datetime FROM klines + WHERE market = ? AND code = ? AND period = ? + AND datetime >= ? + """, + [mkt, code, period, src["_dt"].min().to_pydatetime()], + ).fetchall() + existing_keys = {row[0] for row in existing} + + rows = [] + for _, r in src.iterrows(): + ts = r["_dt"] + rows.append( + ( + mkt, + code, + period, + ts.to_pydatetime(), + float(r["open"]), + float(r["high"]), + float(r["low"]), + float(r["close"]), + float(r["vol"]), + float(r["amount"]), + _row_status(ts), + now, + ) + ) + self._conn.executemany( + """ + INSERT INTO klines (market, code, period, datetime, open, high, low, + close, vol, amount, status, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT (market, code, period, datetime) DO UPDATE SET + open = excluded.open, high = excluded.high, low = excluded.low, + close = excluded.close, vol = excluded.vol, amount = excluded.amount, + status = excluded.status, updated_at = excluded.updated_at + """, + rows, + ) + # DuckDB timestamp 精度:统一到秒级比较 + existing_sec = {ts.replace(microsecond=0) for ts in existing_keys} + inserted = sum(1 for r in rows if r[3].replace(microsecond=0) not in existing_sec) + updated = len(rows) - inserted + return (inserted, updated) + + def promote_provisional(self) -> int: + """把「日期已过」的 provisional 行转正(收盘后的临时值已被次日增量覆盖)。""" + today = datetime.now().date() + cur = self._conn.execute( + """ + UPDATE klines SET status = 'completed' + WHERE status = 'provisional' AND CAST(datetime AS DATE) < ? + """, + [today], + ) + return int(cur.fetchone()[0]) if cur.description else 0 + + # ── 查询 ───────────────────────────────────────────────────────────────── + + def query( + self, + market: str, + code: str, + period: str = "DAILY", + start: str | None = None, + end: str | None = None, + count: int | None = None, + include_provisional: bool = False, + ) -> pd.DataFrame: + """查询 K 线(时间升序),返回 pandas DataFrame。 + + Args: + start / end: 日期范围(``YYYY-MM-DD``,闭区间,可选)。 + count: 只取最近 N 根(在过滤后应用)。 + include_provisional: False(默认)时忽略未收盘的临时 bar—— + 筛选/回测默认口径,杜绝拿盘中价当收盘价。 + """ + mkt = market.upper() + conds = ["market = ?", "code = ?", "period = ?"] + params: list[Any] = [mkt, code, period] + if not include_provisional: + conds.append("status = 'completed'") + if start: + conds.append("CAST(datetime AS DATE) >= ?") + params.append(start) + if end: + conds.append("CAST(datetime AS DATE) <= ?") + params.append(end) + where = " AND ".join(conds) + if count is not None: + df = self._conn.execute( + f""" + SELECT * FROM ( + SELECT * FROM klines WHERE {where} + ORDER BY datetime DESC LIMIT ? + ) ORDER BY datetime ASC + """, + [*params, int(count)], + ).df() + else: + df = self._conn.execute( + f"SELECT * FROM klines WHERE {where} ORDER BY datetime ASC", params + ).df() + return df + + def last_datetime( + self, + market: str, + code: str, + period: str = "DAILY", + include_provisional: bool = True, + ) -> pd.Timestamp | None: + """该标的最新 bar 时间(无数据返回 None)。""" + mkt = market.upper() + cond = "" if include_provisional else "AND status = 'completed'" + row = self._conn.execute( + f""" + SELECT max(datetime) FROM klines + WHERE market = ? AND code = ? AND period = ? {cond} + """, + [mkt, code, period], + ).fetchone() + val = row[0] if row else None + return pd.Timestamp(val) if val is not None else None + + def symbols(self, period: str = "DAILY") -> pd.DataFrame: + """列出仓库内全部标的及其数据范围/行数。""" + return self._conn.execute( + """ + SELECT market, code, count(*) AS bars, + min(datetime) AS first, max(datetime) AS last, + sum(CASE WHEN status = 'provisional' THEN 1 ELSE 0 END) AS provisional + FROM klines WHERE period = ? + GROUP BY market, code ORDER BY market, code + """, + [period], + ).df() + + # ── 删除 ───────────────────────────────────────────────────────────────── + + def delete_symbol(self, market: str, code: str, period: str = "DAILY") -> int: + """删除某标的全部数据,返回删除行数。""" + cur = self._conn.execute( + "DELETE FROM klines WHERE market = ? AND code = ? AND period = ?", + [market.upper(), code, period], + ) + n = cur.fetchone()[0] if cur.description else 0 + return int(n) + + # ── 健康自检(P2-3)────────────────────────────────────────────────────── + + def health_check( + self, + market: str | None = None, + code: str | None = None, + max_gap_weekdays: int = 5, + ) -> dict[str, Any]: + """仓库健康自检:缺口 / 异常跳变 / 最新度 / 临时行统计。 + + Args: + market / code: 只检查指定标的(None = 全仓库)。 + max_gap_weekdays: 相邻 bar 间缺失「交易日数」超过该值报疑似缺口 + (节假日会造成少量误报,报告口径为「待人工核查」)。 + + Returns: + ``{"symbols_checked", "issues": [...], "summary": {...}}``。 + """ + from easy_tdx.mac.qfq_check import detect_ex_dividend_gaps + + conds = ["period = 'DAILY'"] + params: list[Any] = [] + if market is not None: + conds.append("market = ?") + params.append(market.upper()) + if code is not None: + conds.append("code = ?") + params.append(code) + where = " AND ".join(conds) + + sym_df = self._conn.execute( + f""" + SELECT market, code, count(*) AS bars, + min(datetime) AS first, max(datetime) AS last, + sum(CASE WHEN status = 'provisional' THEN 1 ELSE 0 END) AS provisional + FROM klines WHERE {where} + GROUP BY market, code + """, + params, + ).df() + + issues: list[dict[str, Any]] = [] + today = date.today() + stale: list[dict[str, Any]] = [] + total_provisional = 0 + + for _, s in sym_df.iterrows(): + mkt, cde = str(s["market"]), str(s["code"]) + total_provisional += int(s["provisional"]) + + # 最新度 + last = pd.Timestamp(s["last"]) + stale_days = (today - last.date()).days + if stale_days > 7: + stale.append( + {"symbol": f"{mkt}:{cde}", "last": str(last.date()), "days": stale_days} + ) + + # 缺口:相邻 bar 的「工作日差」 + df = self.query(mkt, cde, include_provisional=True) + if len(df) < 2: + continue + dts = pd.to_datetime(df["datetime"]).dt.date + gaps = [] + for prev, cur_d in zip(dts.iloc[:-1], dts.iloc[1:], strict=False): + weekdays = np_busdays(prev, cur_d) + if weekdays > max_gap_weekdays: + gaps.append( + {"after": str(prev), "before": str(cur_d), "weekdays": int(weekdays)} + ) + if gaps: + issues.append( + { + "symbol": f"{mkt}:{cde}", + "kind": "gap", + "detail": ( + f"{len(gaps)} 处疑似缺口(缺失工作日>{max_gap_weekdays},含节假日误报)" + ), + "gaps": gaps[:10], + } + ) + + # 异常跳变(板块感知阈值,复用 QFQ 对拍的跳空检测) + jump_dates = detect_ex_dividend_gaps(df, cde, MARKET_TO_TDX.get(mkt)) + if jump_dates: + issues.append( + { + "symbol": f"{mkt}:{cde}", + "kind": "price_jump", + "detail": ( + f"{len(jump_dates)} 处超出跌停幅度的向下跳空" + "(多为除权,未复权数据需人工核查)" + ), + "dates": jump_dates[:10], + } + ) + + return { + "symbols_checked": int(len(sym_df)), + "issues": issues, + "summary": { + "symbols_with_issues": len({i["symbol"] for i in issues}), + "stale_symbols": stale[:20], + "provisional_rows": total_provisional, + "checked_at": datetime.now().isoformat(timespec="seconds"), + }, + } + + +def np_busdays(d1: date, d2: date) -> int: + """两个日期之间的工作日数(不含首日,含规则近似:周一~周五)。""" + import numpy as _np + + if d2 <= d1: + return 0 + return int(_np.busday_count(d1, d2)) diff --git a/src/easy_tdx/warehouse/sync.py b/src/easy_tdx/warehouse/sync.py new file mode 100644 index 0000000..324fd72 --- /dev/null +++ b/src/easy_tdx/warehouse/sync.py @@ -0,0 +1,155 @@ +"""仓库增量同步器(v1.26 新增)。 + +从 MAC 客户端拉取行情补入 :class:`~easy_tdx.warehouse.store.KlineWarehouse`: + +- **首同步全量**:仓库无该标的数据时按 ``max_bars``(默认 8000 根)拉取; +- **增量补缺**:已有数据时只拉最近 ``tail_bars``(默认 15 根)覆盖—— + 覆盖同日 bar(收盘价修正 / provisional 转正),不动更早历史; +- 同步前自动 :meth:`promote_provisional`(过期临时行转正)。 + +客户端只需具备 ``get_stock_kline(market:int, code, period, start, count, +adjust)`` 签名(``MacClient`` / ``AsyncMacClient`` 均可,本同步器只用同步 +调用——CLI 在主线程使用;serve 的 async 环境请在后台线程调用)。 +""" + +from __future__ import annotations + +import logging +from collections.abc import Callable +from typing import Any + +from easy_tdx.warehouse.store import MARKET_TO_TDX, KlineWarehouse + +logger = logging.getLogger(__name__) + +__all__ = ["WarehouseSyncer", "SyncSummary"] + + +class SyncSummary(dict[str, Any]): + """批次同步结果(dict 语义,键见 :meth:`WarehouseSyncer.sync`)。""" + + +def _period_name(period: str) -> str: + """归一化周期名为仓库 period 键(与客户端 Period 枚举名对齐)。""" + return period.upper() + + +class WarehouseSyncer: + """把客户端行情增量同步进仓库。 + + Example:: + with get_mac_client() as client: + syncer = WarehouseSyncer(client, warehouse) + summary = syncer.sync(["SH:600519", "SZ:000001"]) + print(summary["added"], summary["updated"]) + """ + + def __init__( + self, + client: Any, + warehouse: KlineWarehouse, + max_bars: int = 8000, + tail_bars: int = 15, + adjust: str = "QFQ", + ) -> None: + """Initialize. + + Args: + client: 行情客户端(需有 ``get_stock_kline``)。 + warehouse: 目标仓库。 + max_bars: 首同步(仓库为空时)的最大拉取根数。 + tail_bars: 增量同步拉取的尾部根数(覆盖近几日的修正)。 + adjust: 复权方式(默认 ``"QFQ"`` 前复权——回测/筛选口径; + 注意仓库按 period+adjust 隐含统一,混存不同复权口径需 + 分开仓库文件)。 + """ + self._client = client + self._wh = warehouse + self._max_bars = max(int(max_bars), 800) + self._tail_bars = max(int(tail_bars), 5) + self._adjust = adjust + + def sync_symbol( + self, + market: str, + code: str, + period: str = "DAILY", + ) -> dict[str, Any]: + """同步单个标的,返回 ``{symbol, added, updated, skipped, error}``。""" + symbol = f"{market}:{code}" + try: + existing_last = self._wh.last_datetime(market, code, period) + count = self._tail_bars if existing_last is not None else self._max_bars + df = self._client.get_stock_kline( + MARKET_TO_TDX[market.upper()], + code, + period=period, + start=0, + count=count, + adjust=self._adjust, + ) + if df is None or len(df) == 0: + return {"symbol": symbol, "added": 0, "updated": 0, "skipped": 1, "error": None} + added, updated = self._wh.upsert_bars(market, code, df, period=period) + return { + "symbol": symbol, + "added": added, + "updated": updated, + "skipped": 0, + "error": None, + } + except Exception as exc: # noqa: BLE001 — 单标失败不中断批次 + logger.warning("仓库同步 %s 失败:%s", symbol, exc) + return {"symbol": symbol, "added": 0, "updated": 0, "skipped": 1, "error": str(exc)} + + def sync( + self, + symbols: list[str] | list[tuple[str, str]], + period: str = "DAILY", + progress: Callable[[int, int, str], None] | None = None, + ) -> SyncSummary: + """批量同步,返回汇总。 + + Args: + symbols: ``["SH:600519", ...]`` 或 ``[("SH", "600519"), ...]``。 + period: K 线周期(默认日线)。 + progress: 进度回调 ``progress(done, total, symbol)``。 + + Returns: + ``{"total", "ok", "added", "updated", "skipped", "failed", "details"}``。 + """ + self._wh.promote_provisional() + p = _period_name(period) + + parsed: list[tuple[str, str]] = [] + for s in symbols: + if isinstance(s, str): + mkt, cde = s.split(":", 1) + parsed.append((mkt.strip().upper(), cde.strip())) + else: + parsed.append((s[0].upper(), s[1])) + + details: list[dict[str, Any]] = [] + added = updated = skipped = failed = 0 + for i, (mkt, cde) in enumerate(parsed, 1): + if progress is not None: + progress(i, len(parsed), f"{mkt}:{cde}") + r = self.sync_symbol(mkt, cde, p) + details.append(r) + added += r["added"] + updated += r["updated"] + if r["error"]: + failed += 1 + elif r["skipped"]: + skipped += 1 + return SyncSummary( + { + "total": len(parsed), + "ok": len(parsed) - failed - skipped, + "added": added, + "updated": updated, + "skipped": skipped, + "failed": failed, + "details": details, + } + ) diff --git a/tests/unit/test_warehouse.py b/tests/unit/test_warehouse.py new file mode 100644 index 0000000..348d297 --- /dev/null +++ b/tests/unit/test_warehouse.py @@ -0,0 +1,307 @@ +"""本地 K 线仓库测试(DuckDB store + 增量 sync + provisional 状态机 + 健康自检)。""" + +from __future__ import annotations + +import numpy as np +import pandas as pd +import pytest + +pytest.importorskip("duckdb") + +from easy_tdx.warehouse.store import KlineWarehouse # noqa: E402 +from easy_tdx.warehouse.sync import WarehouseSyncer # noqa: E402 + + +@pytest.fixture() +def wh(tmp_path): + warehouse = KlineWarehouse(tmp_path / "test.duckdb") + yield warehouse + warehouse.close() + + +def _bars(n: int = 10, start: str = "2024-01-01", base: float = 10.0) -> pd.DataFrame: + dates = pd.date_range(start, periods=n, freq="B") + close = base + np.linspace(0, 1, n) + return pd.DataFrame( + { + "datetime": dates, + "open": close, + "high": close * 1.01, + "low": close * 0.99, + "close": close, + "vol": 1000.0, + "amount": close * 1000, + } + ) + + +class _FakeClient: + """返回预置 K 线的假客户端(duck-typed get_stock_kline)。""" + + def __init__(self, df: pd.DataFrame) -> None: + self._df = df + self.calls: list[dict] = [] + + def get_stock_kline(self, market, code, period="DAILY", start=0, count=800, adjust="NONE"): + self.calls.append({"market": market, "count": count, "adjust": adjust}) + return self._df.iloc[max(0, len(self._df) - count) :].reset_index(drop=True) + + +# ── store:写入 / 查询 ─────────────────────────────────────────────────────── + + +def test_upsert_and_query_roundtrip(wh): + df = _bars(10) + added, updated = wh.upsert_bars("SH", "600519", df) + assert (added, updated) == (10, 0) + + out = wh.query("SH", "600519") + assert len(out) == 10 + assert list(out.columns)[:5] == ["market", "code", "period", "datetime", "open"] + assert out["market"].iloc[0] == "SH" + # 升序 + dts = pd.to_datetime(out["datetime"]) + assert dts.is_monotonic_increasing + + +def test_upsert_same_bars_updates_not_duplicates(wh): + df = _bars(10) + wh.upsert_bars("SH", "600519", df) + # 同一批再写 → 全部 update,无重复行 + added, updated = wh.upsert_bars("SH", "600519", df) + assert (added, updated) == (0, 10) + assert len(wh.query("SH", "600519")) == 10 + + +def test_query_count_takes_latest(wh): + full = _bars(50) + wh.upsert_bars("SZ", "000001", full) + out = wh.query("SZ", "000001", count=10) + assert len(out) == 10 + # 是最近 10 根(时间仍升序,且末根 = 全量末根) + last_full = pd.Timestamp(full["datetime"].iloc[-1]).normalize() + assert pd.Timestamp(out["datetime"].iloc[-1]) == last_full + + +def test_query_date_range_filter(wh): + wh.upsert_bars("SZ", "000001", _bars(50)) + out = wh.query("SZ", "000001", start="2024-01-15", end="2024-01-25") + dts = pd.to_datetime(out["datetime"]).dt.date.astype(str) + assert (dts >= "2024-01-15").all() and (dts <= "2024-01-25").all() + + +def test_last_datetime_and_symbols(wh): + assert wh.last_datetime("SH", "600519") is None + wh.upsert_bars("SH", "600519", _bars(10)) + wh.upsert_bars("SZ", "000001", _bars(5)) + assert wh.last_datetime("SH", "600519") == pd.Timestamp("2024-01-12") + syms = wh.symbols() + assert len(syms) == 2 + assert set(syms["code"]) == {"600519", "000001"} + + +def test_delete_symbol(wh): + wh.upsert_bars("SH", "600519", _bars(10)) + assert wh.delete_symbol("SH", "600519") == 10 + assert len(wh.query("SH", "600519")) == 0 + + +def test_missing_optional_columns_filled(wh): + df = _bars(5).drop(columns=["amount"]) + wh.upsert_bars("SH", "600519", df) + out = wh.query("SH", "600519") + assert out["amount"].isna().all() + + +# ── provisional 状态机 ─────────────────────────────────────────────────────── + + +def test_today_bars_before_close_marked_provisional(wh, monkeypatch): + """当日 bar 在 15:05 前落盘 → provisional(逐行判定),默认查询忽略。""" + import datetime as _dt + + import easy_tdx.warehouse.store as store_mod + + today = pd.Timestamp.today().normalize() + dates = pd.date_range(today - pd.Timedelta(days=10), periods=11, freq="D") + df = pd.DataFrame( + { + "datetime": dates, + "open": 10.0, + "high": 10.1, + "low": 9.9, + "close": 10.0, + "vol": 100.0, + "amount": 1000.0, + } + ) + + class _FixedDT(_dt.datetime): + @classmethod + def now(cls, tz=None): # 固定在当日 10:00(盘中) + return _dt.datetime(today.year, today.month, today.day, 10, 0) + + monkeypatch.setattr(store_mod, "datetime", _FixedDT) + + added, _ = wh.upsert_bars("SH", "600519", df) + assert added == 11 + all_rows = wh.query("SH", "600519", include_provisional=True) + completed = wh.query("SH", "600519") + assert len(all_rows) == 11 + assert len(completed) == 10 # 仅当日 bar 是 provisional + # 显式 include_provisional 时当日可见且标记正确 + today_rows = all_rows[all_rows["status"] == "provisional"] + assert len(today_rows) == 1 + assert pd.Timestamp(today_rows["datetime"].iloc[0]).date() == today.date() + + +def test_promote_provisional(wh): + """过期的 provisional 行(日期 < 今天)转正。""" + + old = pd.DataFrame( + { + "datetime": pd.date_range("2024-01-01", periods=3), + "open": 10.0, + "high": 10.1, + "low": 9.9, + "close": 10.0, + "vol": 100.0, + "amount": 1000.0, + } + ) + wh.upsert_bars("SH", "600519", old, status="provisional") + assert len(wh.query("SH", "600519")) == 0 # provisional 默认不可见 + n = wh.promote_provisional() + assert n >= 3 + assert len(wh.query("SH", "600519")) == 3 # 转正后可见 + + +# ── 健康自检 ───────────────────────────────────────────────────────────────── + + +def test_health_check_detects_gap_and_stale(wh): + # 构造缺口:跳过 2 周 + df1 = _bars(5, start="2024-01-01") + df2 = _bars(5, start="2024-03-01") + wh.upsert_bars("SH", "600519", pd.concat([df1, df2], ignore_index=True)) + report = wh.health_check() + assert report["symbols_checked"] == 1 + kinds = [i["kind"] for i in report["issues"]] + assert "gap" in kinds # 1 月→3 月的缺口 + assert report["summary"]["stale_symbols"] # 2024 年数据 → 明显过期 + + +def test_health_check_price_jump(wh): + """除权式跳空被检出(kind=price_jump)。""" + closes = [10.0] * 10 + [7.0] * 10 + dates = pd.date_range("2024-01-01", periods=20, freq="B") + df = pd.DataFrame( + { + "datetime": dates, + "open": closes, + "high": [c * 1.01 for c in closes], + "low": [c * 0.99 for c in closes], + "close": closes, + "vol": 100.0, + "amount": 1000.0, + } + ) + wh.upsert_bars("SH", "600519", df) + report = wh.health_check(market="SH", code="600519") + assert any(i["kind"] == "price_jump" for i in report["issues"]) + + +def test_health_check_clean_series_no_issues(wh): + """连续无跳空数据(工作日)→ 无 gap/price_jump 问题。""" + wh.upsert_bars("SZ", "000001", _bars(30)) + report = wh.health_check(market="SZ", code="000001") + assert report["issues"] == [] + + +# ── 增量同步 ───────────────────────────────────────────────────────────────── + + +def test_sync_initial_full_then_incremental(tmp_path): + warehouse = KlineWarehouse(tmp_path / "s.duckdb") + try: + full = _bars(100) + client = _FakeClient(full) + syncer = WarehouseSyncer(client, warehouse, max_bars=800, tail_bars=15) + + s1 = syncer.sync(["SH:600519"]) + assert s1["added"] == 100 and s1["failed"] == 0 + # 首同步请求了全量(count=800) + assert client.calls[-1]["count"] == 800 + + s2 = syncer.sync([("SH", "600519")]) + assert s2["added"] == 0 and s2["updated"] == 15 # 增量只补尾部 15 根 + assert client.calls[-1]["count"] == 15 + assert len(warehouse.query("SH", "600519")) == 100 # 无重复 + finally: + warehouse.close() + + +def test_sync_new_bars_appended(tmp_path): + warehouse = KlineWarehouse(tmp_path / "s2.duckdb") + try: + client = _FakeClient(_bars(50)) + syncer = WarehouseSyncer(client, warehouse, tail_bars=20) + syncer.sync(["SZ:000001"]) + + # 行情前滚 5 根:新 bar 接在原末根之后 + end = pd.Timestamp(client._df["datetime"].iloc[-1]) + client._df = pd.concat( + [client._df, _bars(5, start=str(end + pd.Timedelta(days=1)))], ignore_index=True + ) + s2 = syncer.sync(["SZ:000001"]) + assert s2["added"] == 5 + assert len(warehouse.query("SZ", "000001")) == 55 + finally: + warehouse.close() + + +def test_sync_failure_does_not_break_batch(tmp_path): + warehouse = KlineWarehouse(tmp_path / "s3.duckdb") + try: + + class _BadClient: + def get_stock_kline(self, *a, **kw): + raise ConnectionError("网络故障") + + syncer = WarehouseSyncer(_BadClient(), warehouse) + s = syncer.sync(["SH:600519", "SZ:000001"]) + assert s["failed"] == 2 + assert all(d["error"] for d in s["details"]) + finally: + warehouse.close() + + +def test_sync_progress_callback(tmp_path): + warehouse = KlineWarehouse(tmp_path / "s4.duckdb") + try: + client = _FakeClient(_bars(20)) + seen: list[tuple[int, int, str]] = [] + + def progress(done, total, sym): + seen.append((done, total, sym)) + + WarehouseSyncer(client, warehouse).sync(["SH:600519", "SZ:000001"], progress=progress) + assert seen == [(1, 2, "SH:600519"), (2, 2, "SZ:000001")] + finally: + warehouse.close() + + +def test_missing_duckdb_helpful_error(tmp_path, monkeypatch): + """duckdb 未安装时给出安装指引(模拟 ImportError)。""" + import builtins + + real_import = builtins.__import__ + + def _no_duckdb(name, *args, **kwargs): + if name == "duckdb": + raise ImportError("No module named 'duckdb'") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _no_duckdb) + with pytest.raises(ImportError, match=r"easy-tdx\[warehouse\]"): + KlineWarehouse(tmp_path / "x.duckdb")