diff --git a/scripts/data_platform/import_vnpy_daily.py b/scripts/data_platform/import_vnpy_daily.py index 8fb70be..fad8f1d 100644 --- a/scripts/data_platform/import_vnpy_daily.py +++ b/scripts/data_platform/import_vnpy_daily.py @@ -11,8 +11,10 @@ import sys import time from pathlib import Path -DB_PATH = os.environ.get('VNPY_DB_PATH', '/tmp/quant_trading_import.db') -DAILY_DIR = '/Volumes/stock/A股数据/日线数据/daily/' +DB_PATH = os.environ.get('VNPY_DB_PATH', '/volume1/stock/sanguo_vnpy/data/quant_trading.db') +# 默认 raw(真实价)—— vnpy CTA 回测撮合用真实价;v1 daily/ 口径已停。 +# 灌 qfq 改 DAILY_DIR=.../qfq/。 +DAILY_DIR = os.environ.get('DAILY_DIR', '/volume1/stock/A股数据/日线数据/raw/') BATCH_SIZE = 50000 # 每批插入行数 diff --git a/scripts/data_platform/merge_increment.py b/scripts/data_platform/merge_increment.py new file mode 100644 index 0000000..603a6b5 --- /dev/null +++ b/scripts/data_platform/merge_increment.py @@ -0,0 +1,191 @@ +#!/usr/bin/env python3 +"""合并 staging 增量到主库(task: 截断 bug 修复)。 + +**核心不变量(截断 bug 回归测试核心)**: + 合并后 main 的行数 >= 合并前 main 的行数(绝不截断)。 + +raw_redownload.py 的 save_one() 用"本次拉的几天"覆盖写整年文件 → 整年数据被截 +(NAS 2026 从 121 行截成 4 行就是这 bug)。本脚本负责按 symbol-year 安全合并: + pd.concat([main, staging]).drop_duplicates(subset=['date'], keep='last') +staging 的新数据/修订值优先(keep='last'),main 已有的历史保留。 + +用法: + python3 merge_increment.py --staging data_cache/daily_update/raw --main data_cache/raw + python3 merge_increment.py --staging ... --main ... --dry-run + +约定: +- staging 与 main 同构:`{root}/{year}/{sh|sz}{code}_daily.parquet` +- staging 文件不修改/不删除(保留供排查),只写 main +""" +from __future__ import annotations + +import argparse +import json +import logging +import os +import sys + +import pandas as pd + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") +log = logging.getLogger("merge_inc") + +# drop_duplicates 的去重键——日线按 date 唯一 +DEDUP_KEY = ["date"] + + +def symbol_from_filename(fname: str) -> str: + """`sh600000_daily.parquet` → `600000`(剥 exchange 前缀和 _daily.parquet 后缀)。""" + base = fname + if base.endswith("_daily.parquet"): + base = base[: -len("_daily.parquet")] + if base[:2] in ("sh", "sz", "bj"): + base = base[2:] + return base + + +def merge_one(staging_file: str, main_file: str, dry_run: bool = False) -> dict: + """合并一个 staging parquet 到 main。 + + 返回 stats: {symbol, before, after, added, action}。 + 若 main 不存在 → action="created"(直接搬 staging 过去)。 + 若 main 存在 → action="merged",按 DEDUP_KEY 去重,staging 值优先(keep='last')。 + """ + sym = symbol_from_filename(os.path.basename(staging_file)) + staging_df = pd.read_parquet(staging_file) + + if not os.path.exists(main_file): + if not dry_run: + os.makedirs(os.path.dirname(main_file), exist_ok=True) + staging_df.to_parquet(main_file, index=False) + return {"symbol": sym, "before": 0, "after": len(staging_df), + "added": len(staging_df), "action": "created"} + + main_df = pd.read_parquet(main_file) + before = len(main_df) + + # concat → drop_duplicates(keep='last' 让 staging 的新/修订值覆盖 main) → sort + merged = pd.concat([main_df, staging_df], ignore_index=True) + merged = merged.drop_duplicates(subset=DEDUP_KEY, keep="last") + merged = merged.sort_values(by="date").reset_index(drop=True) + after = len(merged) + + # 核心不变量:绝不能让 main 变少 + if after < before: + raise RuntimeError( + f"INVARIANT VIOLATED: {main_file} {before}→{after} (staging={staging_file})" + ) + + if not dry_run: + merged.to_parquet(main_file, index=False) + + return {"symbol": sym, "before": before, "after": after, + "added": after - before, "action": "merged"} + + +def walk_staging(staging_root: str) -> list[tuple[str, str]]: + """收集 staging 下所有 `{year}/{sym}_daily.parquet`,返回 [(staging_file, main_file_relpath)]。 + + main_file_relpath 是相对 staging_root 的路径(如 `2026/sh600000_daily.parquet`), + 拼到 main_root 即得 main 文件全路径,保持两边同构。 + """ + pairs: list[tuple[str, str]] = [] + for year in sorted(os.listdir(staging_root)): + ydir = os.path.join(staging_root, year) + if not os.path.isdir(ydir): + continue + for fname in sorted(os.listdir(ydir)): + if not fname.endswith("_daily.parquet"): + continue + rel = os.path.join(year, fname) + pairs.append((os.path.join(ydir, fname), rel)) + return pairs + + +def run_merge(staging_root: str, main_root: str, dry_run: bool = False) -> dict: + """合并 staging → main,返回汇总统计。""" + if not os.path.isdir(staging_root): + raise FileNotFoundError(f"staging dir 不存在: {staging_root}") + + pairs = walk_staging(staging_root) + if not pairs: + log.warning("staging 无 parquet: %s", staging_root) + return {"merged": 0, "created": 0, "skipped": 0, "total_new_rows": 0, + "details": [], "dry_run": dry_run, "ok": True, + "invariant_violations": []} + + os.makedirs(main_root, exist_ok=True) + details: list[dict] = [] + merged_n = created_n = skipped_n = total_new = 0 + invariant_violations: list[str] = [] + + for i, (staging_file, rel) in enumerate(pairs, 1): + main_file = os.path.join(main_root, rel) + try: + stat = merge_one(staging_file, main_file, dry_run=dry_run) + except RuntimeError as e: + # 不变式违反:立刻停(绝不能继续写入更小的 main) + invariant_violations.append(str(e)) + log.error("[%d/%d] INVARIANT %s", i, len(pairs), e) + continue + except Exception as e: # noqa: BLE001 + log.error("[%d/%d] %s ERR %s", i, len(pairs), rel, e) + skipped_n += 1 + continue + + details.append(stat) + if stat["action"] == "created": + created_n += 1 + else: + merged_n += 1 + total_new += stat["added"] + + if i % 1000 == 0 or i == len(pairs): + log.info("[%d/%d] %s %s +%d (main %d→%d)", + i, len(pairs), stat["symbol"], stat["action"], + stat["added"], stat["before"], stat["after"]) + + summary = { + "merged": merged_n, + "created": created_n, + "skipped": skipped_n, + "total_new_rows": total_new, + "details": details, + "dry_run": dry_run, + "invariant_violations": invariant_violations, + } + if invariant_violations: + # 不变式违反致命——即使个别合并成功也判失败 + summary["ok"] = False + else: + summary["ok"] = True + return summary + + +def main(): + ap = argparse.ArgumentParser(description="合并 staging 增量到主库(绝不截断)") + ap.add_argument("--staging", required=True, help="staging 根(如 data_cache/daily_update/raw)") + ap.add_argument("--main", required=True, help="主库根(如 data_cache/raw)") + ap.add_argument("--dry-run", action="store_true", help="只报不写") + ap.add_argument("--summary-json", default=None, help="把汇总写到该 JSON 文件") + args = ap.parse_args() + + summary = run_merge(args.staging, args.main, dry_run=args.dry_run) + + mode = "[DRY-RUN] " if args.dry_run else "" + log.info("=== %s合并完成: merged=%d created=%d skipped=%d 新增行=%d ok=%s ===", + mode, summary["merged"], summary["created"], summary["skipped"], + summary["total_new_rows"], summary["ok"]) + if summary["invariant_violations"]: + log.error("不变式违反 %d 条(main 被截),详情见上", len(summary["invariant_violations"])) + + if args.summary_json: + with open(args.summary_json, "w") as f: + json.dump(summary, f, ensure_ascii=False, indent=2, default=str) + + # 不变式违反 → 退 1(即便部分成功,也提示 main 可能已损坏需排查) + sys.exit(0 if summary["ok"] else 1) + + +if __name__ == "__main__": + main() diff --git a/scripts/data_platform/raw_redownload.py b/scripts/data_platform/raw_redownload.py index fe78dfd..066c1c8 100644 --- a/scripts/data_platform/raw_redownload.py +++ b/scripts/data_platform/raw_redownload.py @@ -24,8 +24,15 @@ import argparse import csv import logging import os +import socket import sys import time +from collections import deque + +# 进程级 socket 超时(关键):akshare 内部 requests 默认无 timeout, +# 遇源限速/慢响应会无限挂起整进程(实测 35 只后 socket 死等 → hang)。 +# 15s 足够区分正常慢响应与挂死,超时即抛 → download_one 捕获 → fail → 继续下一只。 +socket.setdefaulttimeout(15) # 直连:进程级 unset 代理(用户约束) for _k in ["HTTP_PROXY", "HTTPS_PROXY", "http_proxy", "https_proxy", "ALL_PROXY", "all_proxy"]: @@ -42,12 +49,33 @@ RAW_DIR = os.environ.get("RAW_DIR", "/tmp/stock_dl/A股数据/日线数据/raw") SLEEP = float(os.environ.get("SLEEP", "1.0")) DEFAULT_STOCK_LIST = "/tmp/stock_dl/A股数据/stock_info/stock_basic_info_raw_20260326_113530.csv" +# 断路器(借鉴 15min backfill 自愈范式) +CIRCUIT_BREAKER_WINDOW = 200 # 滚动窗口大小 +CIRCUIT_BREAKER_FAIL_RATE = 0.30 # 触发阈值:失败率 > 30% +# 北交所码新浪不支持(必 fail,从断路器分母扣除,防误触发) +KNOWN_UNSUPPORTED_PREFIX = ("920", "921", "83", "87") + def prefix_for(code: str) -> str: """sh/sz 前缀(与 datareader.guess_exchange 一致)。""" return "sh" if code.startswith(("60", "68", "51", "56", "58")) else "sz" +def check_circuit_breaker(recent_results: list) -> bool: + """检查滚动窗口内失败率是否超阈值(断路器触发判定)。 + + recent_results: [(code, ok: bool), ...] 最近下载结果。 + 北交所已知不支持码从分子分母同时扣除(它们必 fail,计入会误触发)。 + 返回 True 表示应触发断路器。 + """ + counted = [(c, ok) for c, ok in recent_results + if not c.startswith(KNOWN_UNSUPPORTED_PREFIX)] + if len(counted) < CIRCUIT_BREAKER_WINDOW: + return False + fails = sum(1 for _, ok in counted if not ok) + return (fails / len(counted)) > CIRCUIT_BREAKER_FAIL_RATE + + def download_one(ak, code: str, start: str, end: str, adjust: str = ""): """新浪源拉日线(adjust="" raw / "qfq" 前复权),返回 (df, None) 或 (None, err)。""" sym = f"{prefix_for(code)}{code}" @@ -79,15 +107,30 @@ def save_one(code: str, df) -> int: return n -def exists(code: str, start_year: int) -> bool: - """symbol 在 start_year 是否已有 parquet(断点续传)。 +def latest_trading_day(end: str) -> str: + """end 往前最近的工作日(周一~周五)。忽略节假日——增量天天跑会自愈。""" + d = pd.Timestamp(end) + while d.weekday() >= 5: # 5=Sat, 6=Sun + d -= pd.Timedelta(days=1) + return d.strftime("%Y-%m-%d") - 同范围续跑 → skip;扩范围(start 更早)→ 新 start_year 不存在 → 重下全量覆盖。 + +def is_fresh(code: str, end: str, start_year: int) -> bool: + """symbol 的 parquet 数据是否已到最新交易日(增量断点续传)。 + + 文件不存在 / 损坏 / 最新日期 < latest_trading_day(end) → 不 fresh → 重下。 + 旧版 exists() 只查文件在不在,对"文件存在但数据旧"的增量场景失效 + (raw/2026 已有旧 parquet → 全 skip → 补数永远不写入)。 """ pref = prefix_for(code) - return os.path.exists( - os.path.join(RAW_DIR, str(start_year), f"{pref}{code}_daily.parquet") - ) + f = os.path.join(RAW_DIR, str(start_year), f"{pref}{code}_daily.parquet") + if not os.path.exists(f): + return False + try: + maxd = pd.read_parquet(f, columns=["date"])["date"].max() + except Exception: # noqa: BLE001 损坏文件 → 重下 + return False + return pd.Timestamp(maxd) >= pd.Timestamp(latest_trading_day(end)) def load_all_codes(stock_list: str) -> list[str]: @@ -135,13 +178,16 @@ def main(): start_year = int(args.start[:4]) ok = fail = rows = skipped = 0 + circuit_recent: deque = deque(maxlen=CIRCUIT_BREAKER_WINDOW) + circuit_count = 0 # 非北交所码下载计数(断路器窗口) for i, code in enumerate(codes, 1): - if not args.force and exists(code, start_year): + if not args.force and is_fresh(code, end, start_year): skipped += 1 if skipped % 500 == 0: - log.info("[%d/%d] ... skipped %d 已存在", i, len(codes), skipped) + log.info("[%d/%d] ... skipped %d 已最新", i, len(codes), skipped) continue df, err = download_one(ak, code, args.start, end, args.adjust) + download_ok = False if df is None: fail += 1 log.warning("[%d/%d] %s FAIL %s", i, len(codes), code, err) @@ -150,10 +196,22 @@ def main(): n = save_one(code, df) ok += 1 rows += n + download_ok = True log.info("[%d/%d] %s ok %d rows", i, len(codes), code, n) except Exception as e: # noqa: BLE001 fail += 1 log.error("[%d/%d] %s SAVE FAIL %s", i, len(codes), code, e) + # 断路器:北交所码不计入(必 fail 会误触发) + if not code.startswith(KNOWN_UNSUPPORTED_PREFIX): + circuit_recent.append((code, download_ok)) + circuit_count += 1 + if circuit_count % CIRCUIT_BREAKER_WINDOW == 0: + if check_circuit_breaker(list(circuit_recent)): + fails = sum(1 for _, ok2 in circuit_recent if not ok2) + rate = fails / len(circuit_recent) * 100 + log.error("断路器触发:最近 %d 只失败率 %.1f%%,abort", + len(circuit_recent), rate) + sys.exit(3) time.sleep(SLEEP) # 限速(用户约束) log.info("=== 完成: ok=%d skip=%d fail=%d rows=%d,raw_dir=%s ===", ok, skipped, fail, rows, RAW_DIR) diff --git a/scripts/data_platform/run_daily_update.sh b/scripts/data_platform/run_daily_update.sh index a9ae067..db6ebfc 100755 --- a/scripts/data_platform/run_daily_update.sh +++ b/scripts/data_platform/run_daily_update.sh @@ -1,27 +1,49 @@ #!/bin/bash # 每日数据增量(C-S3 实走用):raw 真实价 + qfq 前复权,最近 N 天 → NAS。 # +# **安全流程(截断 bug 修复后)**:staging → 验证 → 合并主库 → rsync 主库 NAS。 +# raw_redownload.py 的 save_one() 用"本次几天"覆盖写整年 → 历史被截。 +# 现在下载只写 staging,verify_increment 把关,merge_increment 合并进主库 +# (按 date 去重 keep='last',绝不截断),最后推主库到 NAS。 +# # 跳过 v1 daily_all_update(新浪源接口坏 KeyError:date,C-S3 不用 daily mixed)。 # raw/qfq 用 raw_redownload.py(akshare 新浪 stock_zh_a_daily,全量 29600 已验证)。 # 15min baostock 增量 = 分期项(默认跳,日内策略落地时加)。 # +# **长任务防休眠**:本脚本被 nohup / harness 后台跑前,先 `caffeinate -i -s &` +# (Mac Mini 空闲睡眠会挂进程,详见 memory/feedback-unattended-tasks-prevent-sleep)。 +# # env: -# DAYS=5 增量回看天数(C-S3 实走需当日,5 天兜底停牌/补缺) +# DAYS=7 增量回看天数(C-S3 实走需当日,7 天兜底停牌/补缺/周末) # STOCK_LIST=... 全市场 csv(默认 data_cache/stock_info) +# SKIP_PROBE=1 跳探针预检(调试用) # SKIP_RAW=1 跳 raw 增量 # SKIP_QFQ=1 跳 qfq 增量 +# SKIP_VERIFY=1 跳 verify(调试用,默认把关) +# SKIP_MERGE=1 跳 merge(调试用) +# SKIP_NAS=1 跳推 NAS(烟测用) set -uo pipefail # 不用 -e:单只失败不退(raw_redownload 内部已容错记 fail) NAS=sanguo-nas ROOT="$(cd "$(dirname "$0")/../.." && pwd)" -LOCAL=${STOCK_MOUNT:-$ROOT/data_cache/daily_update} # 持久目录(不进 /tmp,重启不丢) +MAIN=$ROOT/data_cache # 主库(canonical 全量,绝不被下载直接写) +STAGING=$ROOT/data_cache/daily_update # staging(增量暂存,每次清空重下) SL=${STOCK_LIST:-$ROOT/data_cache/stock_info/stock_basic_info_raw_20260326_113530.csv} SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" -DAYS=${DAYS:-5} +DAYS=${DAYS:-7} START=$(python3 -c "import datetime;print((datetime.date.today()-datetime.timedelta(days=$DAYS)).isoformat())") +START_YEAR=${START:0:4} -mkdir -p "$LOCAL/raw" "$LOCAL/qfq" +# 持久日志:stdout+stderr 同时 tee 到文件(保留实时显示 + 持久化) +LOGDIR="$ROOT/data_cache/daily_update/logs" +mkdir -p "$LOGDIR" +LOGFILE="$LOGDIR/daily_$(date +%Y%m%d_%H%M%S).log" +exec > >(tee -a "$LOGFILE") 2>&1 + +mkdir -p "$STAGING/raw" "$STAGING/qfq" "$MAIN/raw" "$MAIN/qfq" +echo "=== 日志: $LOGFILE ===" echo "=== $(date) 每日增量 start=$START DAYS=$DAYS ===" +echo " MAIN=$MAIN STAGING=$STAGING" cd "$SCRIPT_DIR" # 拉 stock_info(代码列表)若本地缺 @@ -30,22 +52,95 @@ if [ ! -f "$SL" ]; then rsync -az "$NAS:/volume1/stock/A股数据/stock_info/" "$(dirname "$SL")/" || true fi +# P. 探针预检(新浪可用性,省时防傻跑 5h 才发现源挂了) +if [ "${SKIP_PROBE:-0}" != "1" ]; then + echo "=== P. 探针预检:sh600000 最近 7 天 ===" + PROBE_START=$(python3 -c "import datetime;print((datetime.date.today()-datetime.timedelta(days=7)).isoformat())") + PROBE_END=$(python3 -c "import datetime;print(datetime.date.today().isoformat())") + if ! PROBE_START="$PROBE_START" PROBE_END="$PROBE_END" timeout 30 python3 -c ' +import os, warnings +for k in ["HTTP_PROXY","HTTPS_PROXY","http_proxy","https_proxy","ALL_PROXY","all_proxy"]: + os.environ.pop(k, None) +os.environ["NO_PROXY"] = "*" +os.environ["no_proxy"] = "*" +warnings.filterwarnings("ignore") +import akshare as ak +df = ak.stock_zh_a_daily( + symbol="sh600000", + start_date=os.environ["PROBE_START"].replace("-",""), + end_date=os.environ["PROBE_END"].replace("-",""), + adjust="", +) +if df is None or df.empty: + print("PROBE: empty result") + raise SystemExit(1) +print("PROBE OK: %d rows" % len(df)) +'; then + echo "!!! 探针失败:新浪不可用,abort 省时" + exit 2 + fi +fi + +# 0. 清空 staging(每次干净增量——staging 是"本次拉的几天",不能累积) +echo "=== 0. 清空 staging parquet ===" +find "$STAGING/raw" -name '*.parquet' -delete 2>/dev/null || true +find "$STAGING/qfq" -name '*.parquet' -delete 2>/dev/null || true + # 1. raw 增量(C-S3 撮合真实价,adjustflag=3 / akshare adjust="") if [ "${SKIP_RAW:-0}" != "1" ]; then - echo "=== 1. raw 增量(最近 $DAYS 天 → $LOCAL/raw)===" - STOCK_LIST="$SL" RAW_DIR="$LOCAL/raw" SLEEP=0.5 \ - python3 raw_redownload.py --all --start "$START" --adjust "" || echo " raw warn(单只失败已记)" + echo "=== 1. raw 增量(最近 $DAYS 天 → $STAGING/raw)===" + STOCK_LIST="$SL" RAW_DIR="$STAGING/raw" SLEEP=0.5 \ + python3 raw_redownload.py --all --start "$START" --adjust "" + rc=$? + if [ "$rc" -eq 3 ]; then + echo "!!! 断路器触发(新浪限流/故障),abort 整个流程" + exit 3 + elif [ "$rc" -ne 0 ]; then + echo " raw warn(单只失败已记, exit=$rc)" + fi fi # 2. qfq 增量(C-S3 warmup 信号,无除权缺口) if [ "${SKIP_QFQ:-0}" != "1" ]; then - echo "=== 2. qfq 增量(最近 $DAYS 天 → $LOCAL/qfq)===" - STOCK_LIST="$SL" RAW_DIR="$LOCAL/qfq" SLEEP=0.5 \ - python3 raw_redownload.py --all --start "$START" --adjust qfq || echo " qfq warn(单只失败已记)" + echo "=== 2. qfq 增量(最近 $DAYS 天 → $STAGING/qfq)===" + STOCK_LIST="$SL" RAW_DIR="$STAGING/qfq" SLEEP=0.5 \ + python3 raw_redownload.py --all --start "$START" --adjust qfq + rc=$? + if [ "$rc" -eq 3 ]; then + echo "!!! 断路器触发(新浪限流/故障),abort 整个流程" + exit 3 + elif [ "$rc" -ne 0 ]; then + echo " qfq warn(单只失败已记, exit=$rc)" + fi fi -# 3. rsync 推 NAS(raw_dir + qfq_dir) -echo "=== 3. rsync → NAS ===" -rsync -az "$LOCAL/raw/" "$NAS:/volume1/stock/A股数据/日线数据/raw/" || echo " raw push warn" -rsync -az "$LOCAL/qfq/" "$NAS:/volume1/stock/A股数据/日线数据/qfq/" || echo " qfq push warn" +# 3. verify staging(安全闸门:不通过则不合并、不推 NAS,staging 留存供排查) +if [ "${SKIP_VERIFY:-0}" != "1" ]; then + echo "=== 3. verify staging ===" + for KIND in raw qfq; do + [ -d "$STAGING/$KIND/$START_YEAR" ] || { echo " $KIND/$START_YEAR 不存在,跳 verify"; continue; } + if ! python3 verify_increment.py --staging "$STAGING/$KIND" --start "$START"; then + echo " !!! $KIND verify FAILED —— 不合并、不推 NAS,staging 留存排查" + echo " !!! 详见上方 JSON 输出(failed_symbols / fatal_samples)" + exit 1 + fi + done +fi + +# 4. 合并 staging → 主库(按 date 去重 keep='last',绝不截断) +if [ "${SKIP_MERGE:-0}" != "1" ]; then + echo "=== 4. merge staging → main ===" + for KIND in raw qfq; do + [ -d "$STAGING/$KIND/$START_YEAR" ] || { echo " $KIND/$START_YEAR 不存在,跳 merge"; continue; } + python3 merge_increment.py --staging "$STAGING/$KIND" --main "$MAIN/$KIND" \ + || { echo " !!! $KIND merge 失败(可能不变式违反),不推 NAS"; exit 1; } + done +fi + +# 5. rsync 主库 → NAS(改:推 $MAIN 不是 $STAGING;原脚本推 staging 是 bug 之一) +if [ "${SKIP_NAS:-0}" != "1" ]; then + echo "=== 5. rsync 主库 → NAS ===" + rsync -az "$MAIN/raw/" "$NAS:/volume1/stock/A股数据/日线数据/raw/" || echo " raw push warn" + rsync -az "$MAIN/qfq/" "$NAS:/volume1/stock/A股数据/日线数据/qfq/" || echo " qfq push warn" +fi echo "=== 完成 $(date) ===" diff --git a/scripts/data_platform/verify_increment.py b/scripts/data_platform/verify_increment.py new file mode 100644 index 0000000..62bed84 --- /dev/null +++ b/scripts/data_platform/verify_increment.py @@ -0,0 +1,202 @@ +#!/usr/bin/env python3 +"""验证 staging 增量质量(task: 截断 bug 修复的安全闸门)。 + +在 merge_increment 之前跑:校验 staging 的 raw/qfq 增量数据是否合格, +不合格则 run_daily_update.sh 不合并、不推 NAS(staging 留存供排查)。 + +**复用 scripts/data_platform/validator.py 的 DataValidator(七条 fatal)**: + D1 价格>0 / D2 OHLC 一致 / D3 volume≥0 / D6 日期不重复 / D7 非未来日期 / ... + +阈值(Main Agent 已定,硬编码常量): + MIN_SUCCESS_RATE = 0.95 成功率(分母扣除北交所已知不支持码) + MIN_FRESH_RATE = 0.95 最大日期 >= 最近交易日的 symbol 占比 + KNOWN_UNSUPPORTED_PREFIX 新浪 stock_zh_a_daily 不支持的北交所码(从分母扣) + +用法: + python3 verify_increment.py --staging data_cache/daily_update/raw \\ + --stock-list data_cache/stock_info/stock_basic_info_raw_*.csv \\ + --start 2026-07-01 + +退出码:passed=0,否则 1。输出 JSON 到 stdout。 +""" +from __future__ import annotations + +import argparse +import json +import logging +import os +import sys +from datetime import datetime, timedelta + +import pandas as pd + +# 同目录 import(脚本运行目录 = scripts/data_platform/) +_SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) +if _SCRIPT_DIR not in sys.path: + sys.path.insert(0, _SCRIPT_DIR) + +from validator import DataValidator # noqa: E402 + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") +log = logging.getLogger("verify_inc") + +# 阈值(Main Agent 已定) +MIN_SUCCESS_RATE = 0.95 +MIN_FRESH_RATE = 0.95 +KNOWN_UNSUPPORTED_PREFIX = ("920", "921", "83", "87") +MAX_FATAL_SAMPLES = 5 # 输出里只放前 N 个失败样本(避免输出爆掉) + + +def _latest_available_trading_day(now: datetime | None = None) -> str: + """A 股收盘感知:返回 now 时点最近的可获取交易日(YYYY-MM-DD)。 + + catch-up 跨夜场景:下载发生在昨晚、verify 跑在今天盘前, + 今天数据尚未生成(15:00 才收盘)→ 目标应是昨天而非今天。 + + - 工作日 >= 15:00 → 今天(已收盘,数据可获取) + - 工作日 < 15:00 → 上一交易日(今天未收盘,回退到工作日) + - 周末 → 上周五 + + 注意:与 raw_redownload.latest_trading_day 不同——那个只回退周末、 + 用于 is_fresh 决定是否重下;本函数额外感知收盘时间,用于 verify 闸门。 + """ + if now is None: + now = datetime.now() + d = now + if d.weekday() >= 5: # 周末 → 回退到周五 + d = d - timedelta(days=(d.weekday() - 4)) + elif d.hour < 15: # 工作日盘前 → 上一交易日 + d = d - timedelta(days=1) + while d.weekday() >= 5: # 回退周末 + d = d - timedelta(days=1) + return d.strftime("%Y-%m-%d") + + +def symbol_from_filename(fname: str) -> str: + """`sh600000_daily.parquet` → `600000`。""" + base = fname + if base.endswith("_daily.parquet"): + base = base[: -len("_daily.parquet")] + if base[:2] in ("sh", "sz", "bj"): + base = base[2:] + return base + + +def verify(staging_root: str, start: str, _now: datetime | None = None) -> dict: + """校验 staging 下 start_year 目录的所有 parquet。 + + 返回 dict: {passed, success_rate, fresh_rate, total, unsupported_skipped, + success, fresh, failed_symbols, fatal_samples, latest_trading_day}。 + + _now 仅用于测试注入固定时间;生产留空取 datetime.now()。 + """ + start_year = int(start[:4]) + year_dir = os.path.join(staging_root, str(start_year)) + if not os.path.isdir(year_dir): + raise FileNotFoundError(f"staging year dir 不存在: {year_dir}") + + now = _now or datetime.now() + latest = _latest_available_trading_day(now) + latest_ts = pd.Timestamp(latest) + log.info("最近可获取交易日=%s(now=%s)", latest, now.strftime("%Y-%m-%d %H:%M")) + + validator = DataValidator() + + total = 0 + unsupported_skipped = 0 + success = 0 + fresh = 0 + failed_symbols: list[str] = [] + fatal_samples: list[dict] = [] + + files = sorted(f for f in os.listdir(year_dir) if f.endswith("_daily.parquet")) + for fname in files: + total += 1 + code = symbol_from_filename(fname) + # 北交所码新浪不支持(log 大量 920xxx KeyError:'date' 确认)→ 从分母扣 + if code.startswith(KNOWN_UNSUPPORTED_PREFIX): + unsupported_skipped += 1 + continue + + fpath = os.path.join(year_dir, fname) + try: + df = pd.read_parquet(fpath) + except Exception as e: # noqa: BLE001 损坏文件记 fail + failed_symbols.append(code) + if len(fatal_samples) < MAX_FATAL_SAMPLES: + fatal_samples.append({"symbol": code, "errors": [f"read_parquet: {type(e).__name__}: {str(e)[:80]}"]}) + continue + + result = validator.validate(df, "daily") + if not result.passed: + failed_symbols.append(code) + if len(fatal_samples) < MAX_FATAL_SAMPLES: + fatal_samples.append({"symbol": code, "errors": result.fatal_errors[:3]}) + continue + + success += 1 + # 新鲜度:最大日期 >= 最近交易日(staging 是增量,多数会 == latest) + try: + maxd = pd.Timestamp(df["date"].max()) + except Exception: # noqa: BLE001 无 date 列记为不 fresh(但已过校验,理论上不会) + maxd = pd.Timestamp("1970-01-01") + if maxd >= latest_ts: + fresh += 1 + + denom = total - unsupported_skipped + success_rate = (success / denom) if denom > 0 else 0.0 + fresh_rate = (fresh / success) if success > 0 else 0.0 + passed = (success_rate >= MIN_SUCCESS_RATE) and (fresh_rate >= MIN_FRESH_RATE) and (denom > 0) + + return { + "passed": passed, + "success_rate": round(success_rate, 4), + "fresh_rate": round(fresh_rate, 4), + "total": total, + "unsupported_skipped": unsupported_skipped, + "success": success, + "fresh": fresh, + "denom": denom, + "failed_count": len(failed_symbols), + "failed_symbols": failed_symbols, + "fatal_samples": fatal_samples, + "latest_trading_day": latest, + "thresholds": { + "MIN_SUCCESS_RATE": MIN_SUCCESS_RATE, + "MIN_FRESH_RATE": MIN_FRESH_RATE, + "KNOWN_UNSUPPORTED_PREFIX": list(KNOWN_UNSUPPORTED_PREFIX), + }, + } + + +def main(): + ap = argparse.ArgumentParser(description="验证 staging 增量质量") + ap.add_argument("--staging", required=True, help="staging 根(如 data_cache/daily_update/raw)") + ap.add_argument("--stock-list", default=None, help="stock_basic_info csv(保留接口,本版未强用)") + ap.add_argument("--start", required=True, help="增量起点 YYYY-MM-DD(决定看哪个 year 目录)") + ap.add_argument("--summary-json", default=None, help="把结果写到该 JSON 文件") + args = ap.parse_args() + + try: + result = verify(args.staging, args.start) + except FileNotFoundError as e: + log.error("%s", e) + sys.exit(2) + + # 输出 + print(json.dumps(result, ensure_ascii=False, indent=2, default=str)) + + if args.summary_json: + with open(args.summary_json, "w") as f: + json.dump(result, f, ensure_ascii=False, indent=2, default=str) + + verdict = "PASSED" if result["passed"] else "FAILED" + log.info("=== verify %s: success_rate=%.2f fresh_rate=%.2f total=%d skipped=%d failed=%d ===", + verdict, result["success_rate"], result["fresh_rate"], + result["total"], result["unsupported_skipped"], result["failed_count"]) + + sys.exit(0 if result["passed"] else 1) + + +if __name__ == "__main__": + main() diff --git a/tests/data/test_index_downloader.py b/tests/data/test_index_downloader.py new file mode 100644 index 0000000..0888f2a --- /dev/null +++ b/tests/data/test_index_downloader.py @@ -0,0 +1,197 @@ +"""Tests for index downloader and read_index_daily functionality.""" +import pandas as pd +import os +from unittest.mock import patch, MagicMock +from datetime import date +from pathlib import Path +import pytest + +from sanguo_data.config import DataConfig + + +def test_download_index_writes_parquet(tmp_path): + """Test that download_index writes parquet files with correct structure.""" + # Sample data that baostock would return + sample_data = [ + ["2024-01-02", "3495.0", "3505.0", "3490.0", "3500.0", "100000"], + ["2024-01-03", "3505.0", "3515.0", "3500.0", "3510.0", "120000"], + ["2024-01-04", "3515.0", "3525.0", "3510.0", "3520.0", "110000"], + ] + + # Create a simple baostock mock + class MockBaostock: + class MockResult: + def __init__(self, data): + self.error_code = "success" + self.error_msg = "success" + self.data = data + self.fields = ["date", "open", "high", "low", "close", "volume"] + self.row_index = 0 + + def next(self): + if self.row_index < len(self.data): + row = self.data[self.row_index] + self.row_index += 1 + return True + return False + + def get_row_data(self): + return self.data[self.row_index - 1] + + def login(self): + return self.MockResult([]) + + def logout(self): + return self.MockResult([]) + + def query_history_k_data_plus(self, *args, **kwargs): + return self.MockResult(sample_data) + + # Patch baostock module + import sys + sys.modules["baostock"] = MockBaostock() + + try: + # Import after patching + from sanguo_data.index_downloader import download_index + + # Download index data + download_index("sh000300", 2024, 2024, str(tmp_path)) + + finally: + # Clean up the mock + del sys.modules["baostock"] + + # Verify parquet file was created + expected_file = tmp_path / "2024" / "sh000300_daily.parquet" + assert expected_file.exists(), f"Expected parquet file at {expected_file}" + + # Verify parquet content + df_read = pd.read_parquet(expected_file) + assert len(df_read) == 3 + assert "close" in df_read.columns + assert "date" in df_read.columns + assert df_read["close"].iloc[0] == 3500.0 + + +def test_download_index_clears_proxy(tmp_path): + """Test that download_index clears proxy environment variables.""" + # Set proxy variables + os.environ["http_proxy"] = "http://evil:8080" + os.environ["https_proxy"] = "https://evil:8080" + + sample_data = [["2024-01-02", "3495.0", "3505.0", "3490.0", "3500.0", "100000"]] + + # Create a simple baostock mock + class MockBaostock: + class MockResult: + def __init__(self, data): + self.error_code = "success" + self.error_msg = "success" + self.data = data + self.fields = ["date", "open", "high", "low", "close", "volume"] + self.row_index = 0 + + def next(self): + if self.row_index < len(self.data): + row = self.data[self.row_index] + self.row_index += 1 + return True + return False + + def get_row_data(self): + return self.data[self.row_index - 1] + + def login(self): + return self.MockResult([]) + + def logout(self): + return self.MockResult([]) + + def query_history_k_data_plus(self, *args, **kwargs): + return self.MockResult(sample_data) + + # Patch baostock module + import sys + sys.modules["baostock"] = MockBaostock() + + try: + # Import after patching + from sanguo_data.index_downloader import download_index + + # Download index data + download_index("sh000300", 2024, 2024, str(tmp_path)) + + finally: + # Clean up the mock + del sys.modules["baostock"] + + # Verify proxy variables were cleared + assert "http_proxy" not in os.environ + assert "https_proxy" not in os.environ + + +def test_read_index_daily_reads_parquet(tmp_path): + """Test that read_index_daily reads index parquet files correctly.""" + # Import after implementation + from sanguo_data.datareader import read_index_daily + + # Create test parquet file with same structure as stock data + year_dir = tmp_path / "2024" + year_dir.mkdir() + + df = pd.DataFrame({ + "date": ["2024-01-02", "2024-01-03", "2024-01-04"], + "open": [3495.0, 3505.0, 3515.0], + "high": [3505.0, 3515.0, 3525.0], + "low": [3490.0, 3500.0, 3510.0], + "close": [3500.0, 3510.0, 3520.0], + "volume": [100000, 120000, 110000], + }) + df.to_parquet(year_dir / "sh000300_daily.parquet") + + cfg = DataConfig( + data_paths={"daily_dir": str(tmp_path)}, + data_sources={}, validation={}, performance={}, + ) + + # Read index daily data + result = read_index_daily("sh000300", date(2024, 1, 1), date(2024, 12, 31), cfg) + + # Verify result + assert len(result) == 3 + assert "close" in result.columns + assert result["close"].iloc[0] == 3500.0 + + +def test_read_index_daily_handles_date_range(tmp_path): + """Test that read_index_daily filters by date range correctly.""" + # Import after implementation + from sanguo_data.datareader import read_index_daily + + # Create test parquet file + year_dir = tmp_path / "2024" + year_dir.mkdir() + + df = pd.DataFrame({ + "date": ["2024-01-02", "2024-06-15", "2024-12-31"], + "open": [3495.0, 3600.0, 3700.0], + "high": [3505.0, 3610.0, 3710.0], + "low": [3490.0, 3590.0, 3690.0], + "close": [3500.0, 3605.0, 3705.0], + "volume": [100000, 120000, 110000], + }) + df.to_parquet(year_dir / "sh000300_daily.parquet") + + cfg = DataConfig( + data_paths={"daily_dir": str(tmp_path)}, + data_sources={}, validation={}, performance={}, + ) + + # Read with narrowed date range + result = read_index_daily("sh000300", date(2024, 1, 1), date(2024, 6, 30), cfg) + + # Should only get first 2 rows + assert len(result) == 2 + assert result["close"].iloc[0] == 3500.0 + assert result["close"].iloc[1] == 3605.0 diff --git a/tests/data_platform/__init__.py b/tests/data_platform/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/data_platform/test_circuit_breaker.py b/tests/data_platform/test_circuit_breaker.py new file mode 100644 index 0000000..19ec481 --- /dev/null +++ b/tests/data_platform/test_circuit_breaker.py @@ -0,0 +1,155 @@ +"""断路器单元测试(task: Phase 2 可靠性增强)。 + +测试 check_circuit_breaker 纯函数——不启动整个下载流程, +直接构造 recent_results 列表验证触发逻辑。 + +覆盖场景: + 1. 35% 失败率 → 触发 + 2. 25% 失败率 → 不触发 + 3. 北交所 920xxx 的 fail 不计入 → 不触发 + 4. 不足窗口 → 不触发 + 5. 恰好 30% → 不触发(阈值是 > 0.30,不含等于) +""" +import os +import sys + +import pytest + +_SCRIPT_DIR = os.path.join(os.path.dirname(__file__), "..", "..", "scripts", "data_platform") +_SCRIPT_DIR = os.path.abspath(_SCRIPT_DIR) +if _SCRIPT_DIR not in sys.path: + sys.path.insert(0, _SCRIPT_DIR) + +from raw_redownload import ( # noqa: E402 + CIRCUIT_BREAKER_FAIL_RATE, + CIRCUIT_BREAKER_WINDOW, + KNOWN_UNSUPPORTED_PREFIX, + check_circuit_breaker, +) + + +def _make_results(n_ok: int, n_fail: int, fail_prefix: str = "600") -> list: + """构造 recent_results 列表:n_ok 个 ok + n_fail 个 fail。 + + fail_prefix 控制失败码的前缀(用于测试北交所排除)。 + """ + ok_list = [("600001", True)] * n_ok + fail_list = [(f"{fail_prefix}9999", False)] * n_fail + return ok_list + fail_list + + +class TestCircuitBreakerTrigger: + """断路器触发阈值测试。""" + + def test_35_percent_fail_triggers(self): + """200 只里 70 只 fail(35%)→ 触发。""" + results = _make_results(n_ok=130, n_fail=70) + assert len(results) == 200 + assert check_circuit_breaker(results) is True + + def test_25_percent_fail_no_trigger(self): + """200 只里 50 只 fail(25%)→ 不触发。""" + results = _make_results(n_ok=150, n_fail=50) + assert len(results) == 200 + assert check_circuit_breaker(results) is False + + def test_exactly_30_percent_no_trigger(self): + """恰好 30%(60/200)→ 不触发(阈值是 > 0.30,不含等于)。""" + results = _make_results(n_ok=140, n_fail=60) + assert len(results) == 200 + fail_rate = 60 / 200 + assert fail_rate == pytest.approx(CIRCUIT_BREAKER_FAIL_RATE) + assert check_circuit_breaker(results) is False + + def test_31_percent_triggers(self): + """31% 失败率 → 触发(刚过阈值)。""" + results = _make_results(n_ok=138, n_fail=62) + assert len(results) == 200 + assert check_circuit_breaker(results) is True + + +class TestCircuitBreakerBjExclusion: + """北交所码不计入断路器分母测试。""" + + def test_bj_fails_excluded_from_denominator(self): + """北交所 920xxx 的 fail 不计入 → 不触发。 + + 200 只非北交所全部 ok + 100 只北交所全部 fail: + 有效分母=200,有效失败=0 → 0% → 不触发。 + """ + ok_list = [("600001", True)] * 200 + bj_fail_list = [("920001", False)] * 100 + results = ok_list + bj_fail_list + assert check_circuit_breaker(results) is False + + def test_bj_fails_do_not_inflate_rate(self): + """北交所 fail 混在 200 窗口内不抬高失败率。 + + 150 只非北交所 ok + 50 只北交所 fail = 200 只总数: + 有效分母=150(< 窗口 200)→ 不足窗口,不触发。 + """ + ok_list = [("600001", True)] * 150 + bj_fail_list = [("920001", False)] * 50 + results = ok_list + bj_fail_list + assert check_circuit_breaker(results) is False + + def test_mixed_bj_and_normal_fails(self): + """混合场景:200 非北交所(60 fail=30%)+ 50 北交所 fail。 + + 有效:200 非北交所,60 fail = 30%,恰好阈值(> 0.30 不含等于)→ 不触发。 + 北交所 50 fail 被排除,不影响计算。 + """ + ok_normal = [("600001", True)] * 140 + fail_normal = [("600002", False)] * 60 + fail_bj = [("920001", False)] * 50 + results = ok_normal + fail_normal + fail_bj + assert check_circuit_breaker(results) is False + + def test_all_bj_prefixes_excluded(self): + """所有已知不支持前缀都排除:920/921/83/87。""" + ok_list = [("600001", True)] * 200 + bj_fails = ( + [("920001", False)] * 30 + + [("921002", False)] * 30 + + [("830003", False)] * 30 + + [("870004", False)] * 30 + ) + results = ok_list + bj_fails + assert check_circuit_breaker(results) is False + + +class TestCircuitBreakerWindow: + """窗口大小边界测试。""" + + def test_insufficient_window_no_trigger(self): + """不足 200 只 → 不触发(即使全部 fail)。""" + results = [("600001", False)] * 199 + assert check_circuit_breaker(results) is False + + def test_exactly_window_triggers_if_high_fail(self): + """恰好 200 只且失败率超阈值 → 触发。""" + results = _make_results(n_ok=100, n_fail=100) + assert len(results) == 200 + assert check_circuit_breaker(results) is True + + def test_empty_results_no_trigger(self): + """空列表 → 不触发。""" + assert check_circuit_breaker([]) is False + + def test_all_ok_no_trigger(self): + """200 只全部 ok → 不触发。""" + results = [("600001", True)] * 200 + assert check_circuit_breaker(results) is False + + +class TestCircuitBreakerConstants: + """常量值校验(防止意外修改)。""" + + def test_window_is_200(self): + assert CIRCUIT_BREAKER_WINDOW == 200 + + def test_fail_rate_is_030(self): + assert CIRCUIT_BREAKER_FAIL_RATE == 0.30 + + def test_known_unsupported_prefixes(self): + assert KNOWN_UNSUPPORTED_PREFIX == ("920", "921", "83", "87") diff --git a/tests/data_platform/test_merge_increment.py b/tests/data_platform/test_merge_increment.py new file mode 100644 index 0000000..48d97e4 --- /dev/null +++ b/tests/data_platform/test_merge_increment.py @@ -0,0 +1,179 @@ +"""Tests for merge_increment.py — 截断 bug 回归核心。 + +不变量:合并后 main 行数 >= 合并前 main 行数(绝不截断)。 +""" +from __future__ import annotations + +import os +import sys + +import pandas as pd +import pytest + +# 让测试能 import scripts/data_platform/ 下的模块 +_HERE = os.path.dirname(os.path.abspath(__file__)) +_SCRIPT_DIR = os.path.abspath(os.path.join(_HERE, "..", "..", "scripts", "data_platform")) +if _SCRIPT_DIR not in sys.path: + sys.path.insert(0, _SCRIPT_DIR) + +from merge_increment import merge_one, run_merge, symbol_from_filename, walk_staging # noqa: E402 + + +# ---------- helpers ---------- + +def _make_df(dates: list[str], close_base: float = 10.0) -> pd.DataFrame: + """构造合法日线 df(date + OHLCV)。close 自 close_base 递增,便于区分来源。""" + n = len(dates) + return pd.DataFrame({ + "date": pd.to_datetime(dates), + "open": [close_base + i for i in range(n)], + "high": [close_base + i + 0.5 for i in range(n)], + "low": [close_base + i - 0.2 for i in range(n)], + "close": [close_base + i + 0.3 for i in range(n)], + "volume": [10000 + i for i in range(n)], + }) + + +def _make_main_with_100_rows(path: str) -> int: + """在 path 写一份 100 行的 main,返回行数。""" + dates = pd.bdate_range("2026-01-01", periods=100).strftime("%Y-%m-%d").tolist() + df = _make_df(dates, close_base=10.0) + os.makedirs(os.path.dirname(path), exist_ok=True) + df.to_parquet(path, index=False) + return len(df) + + +# ---------- tests ---------- + +def test_symbol_from_filename(): + assert symbol_from_filename("sh600000_daily.parquet") == "600000" + assert symbol_from_filename("sz000001_daily.parquet") == "000001" + assert symbol_from_filename("bj920000_daily.parquet") == "920000" + + +def test_merge_one_new_main(tmp_path): + """main 不存在 → created,行数 = staging 行数。""" + staging = tmp_path / "staging" / "2026" / "sh600000_daily.parquet" + main = tmp_path / "main" / "2026" / "sh600000_daily.parquet" + _make_main_with_100_rows(str(staging)) # 这里 staging 当作源写 + + stat = merge_one(str(staging), str(main)) + assert stat["action"] == "created" + assert stat["before"] == 0 + assert stat["after"] == 100 + assert os.path.exists(main) + + +def test_merge_one_no_truncation_invariant(tmp_path): + """**核心回归**:main=100 + staging=5(3新+2重复) → 合并后 103(>=100,不截断)。""" + staging = tmp_path / "staging" / "2026" / "sh600000_daily.parquet" + main = tmp_path / "main" / "2026" / "sh600000_daily.parquet" + staging.parent.mkdir(parents=True, exist_ok=True) + main.parent.mkdir(parents=True, exist_ok=True) + + # main: 100 行(2026-01-01 起) + main_dates = pd.bdate_range("2026-01-01", periods=100).strftime("%Y-%m-%d").tolist() + _make_df(main_dates, close_base=10.0).to_parquet(main, index=False) + before_rows = 100 + + # staging: 5 行 = 最后 2 个已有日期(重复,用于测 keep='last')+ 3 个新日期 + dup_dates = main_dates[-2:] # 例如 ...0408, 0409 + last_main = pd.Timestamp(main_dates[-1]) + new_dates = pd.bdate_range(last_main + pd.Timedelta(days=1), periods=3).strftime("%Y-%m-%d").tolist() + staging_dates = dup_dates + new_dates + staging_df = _make_df(staging_dates, close_base=999.0) # 999 让 staging 值可识别 + staging_df.to_parquet(staging, index=False) + + stat = merge_one(str(staging), str(main)) + + # 不变式:绝不截断 + assert stat["before"] == before_rows + assert stat["after"] >= before_rows, f"INVARIANT: {stat['before']}→{stat['after']}" + # 精确:100 + 3 新 = 103(2 个重复去重后保留 staging 值) + assert stat["after"] == 103 + assert stat["added"] == 3 + assert stat["action"] == "merged" + + # 重复日期 → keep='last' 取 staging 值(999.x) + merged = pd.read_parquet(main) + dup_row = merged[merged["date"] == pd.Timestamp(dup_dates[0])].iloc[0] + assert dup_row["close"] == pytest.approx(999.3), "重复日期应取 staging 值" + + # 新日期都在 + for d in new_dates: + assert pd.Timestamp(d) in merged["date"].values + + +def test_merge_one_dry_run_no_write(tmp_path): + """dry-run 不写 main。""" + staging = tmp_path / "staging" / "2026" / "sh600000_daily.parquet" + main = tmp_path / "main" / "2026" / "sh600000_daily.parquet" + _make_main_with_100_rows(str(staging)) + # main 不存在,dry-run 应保持不存在 + stat = merge_one(str(staging), str(main), dry_run=True) + assert stat["action"] == "created" + assert not os.path.exists(main) + + +def test_run_merge_summary_and_invariant(tmp_path): + """端到端:多 symbol 合并 + 汇总统计 + 不变式全过。""" + staging_root = tmp_path / "staging" + main_root = tmp_path / "main" + + # 构造 3 只 symbol:2 只 main 已有需合并,1 只 main 没有需 created + setup = [ + ("sh600000", True), # main 存在 + ("sz000001", True), # main 存在 + ("sh600004", False), # main 不存在 + ] + for sym, has_main in setup: + year = "2026" + main_dates = pd.bdate_range("2026-01-01", periods=50).strftime("%Y-%m-%d").tolist() + last_main = pd.Timestamp(main_dates[-1]) + new3 = pd.bdate_range(last_main + pd.Timedelta(days=1), periods=3).strftime("%Y-%m-%d").tolist() + staging_dates = main_dates[-1:] + new3 + sdir = staging_root / year + mdir = main_root / year + sdir.mkdir(parents=True, exist_ok=True) + _make_df(staging_dates, close_base=888.0).to_parquet(sdir / f"{sym}_daily.parquet", index=False) + if has_main: + mdir.mkdir(parents=True, exist_ok=True) + _make_df(main_dates, close_base=10.0).to_parquet(mdir / f"{sym}_daily.parquet", index=False) + + summary = run_merge(str(staging_root), str(main_root)) + + assert summary["ok"] is True + assert summary["merged"] == 2 + assert summary["created"] == 1 + assert summary["skipped"] == 0 + # merged: 每只新增 3(staging 4 - 1 重复);created: 新增 4(全 staging) + assert summary["total_new_rows"] == 10 + # 不变式违反 0 + assert summary["invariant_violations"] == [] + # 不变式:merged 类 main 行数 >= 原 main 大小(50);created 类只 >= staging 大小(4) + for sym, has_main in setup: + df = pd.read_parquet(main_root / "2026" / f"{sym}_daily.parquet") + threshold = 50 if has_main else 4 + assert len(df) >= threshold, f"{sym} main 行数 {len(df)} < {threshold}" + + +def test_run_merge_empty_staging(tmp_path): + """staging 无 parquet → merged=0, ok=True(空也算安全通过)。""" + staging_root = tmp_path / "staging" + staging_root.mkdir() + main_root = tmp_path / "main" + summary = run_merge(str(staging_root), str(main_root)) + assert summary["merged"] == 0 + assert summary["ok"] is True + + +def test_walk_staging_collects_pairs(tmp_path): + """walk_staging 应只收 *.parquet,跳非 parquet 和非目录。""" + s = tmp_path / "staging" + (s / "2026").mkdir(parents=True) + (s / "2026" / "sh600000_daily.parquet").write_bytes(b"x") + (s / "2026" / "README.txt").write_text("nope") + (s / "not_a_year.txt").write_text("nope") + pairs = walk_staging(str(s)) + assert len(pairs) == 1 + assert pairs[0][1] == os.path.join("2026", "sh600000_daily.parquet") diff --git a/tests/data_platform/test_verify_increment.py b/tests/data_platform/test_verify_increment.py new file mode 100644 index 0000000..d1d907f --- /dev/null +++ b/tests/data_platform/test_verify_increment.py @@ -0,0 +1,238 @@ +"""Tests for verify_increment.py — 安全闸门。""" +from __future__ import annotations + +import os +import sys +import time +from datetime import datetime + +import pandas as pd +import pytest + +_HERE = os.path.dirname(os.path.abspath(__file__)) +_SCRIPT_DIR = os.path.abspath(os.path.join(_HERE, "..", "..", "scripts", "data_platform")) +if _SCRIPT_DIR not in sys.path: + sys.path.insert(0, _SCRIPT_DIR) + +import verify_increment as vi # noqa: E402 + + +# ---------- helpers ---------- + +def _good_df(dates: list[str]) -> pd.DataFrame: + """合法日线(通过 DataValidator 所有 fatal)。""" + n = len(dates) + return pd.DataFrame({ + "date": pd.to_datetime(dates), + "open": [10.0 + i for i in range(n)], + "high": [10.5 + i for i in range(n)], + "low": [9.8 + i for i in range(n)], + "close": [10.2 + i for i in range(n)], + "volume": [10000 + i for i in range(n)], + }) + + +def _write_staging(staging_root: str, year: str, fname: str, df: pd.DataFrame) -> None: + ydir = os.path.join(staging_root, year) + os.makedirs(ydir, exist_ok=True) + df.to_parquet(os.path.join(ydir, fname), index=False) + + +def _recent_dates(n: int = 3) -> list[str]: + """最近 n 个工作日(保证 fresh,含今天/最近交易日)。""" + today = pd.Timestamp(time.strftime("%Y-%m-%d")) + dates = pd.bdate_range(end=today, periods=n).strftime("%Y-%m-%d").tolist() + return dates + + +# ---------- symbol_from_filename ---------- + +def test_symbol_from_filename(): + assert vi.symbol_from_filename("sh600000_daily.parquet") == "600000" + assert vi.symbol_from_filename("sz000001_daily.parquet") == "000001" + assert vi.symbol_from_filename("bj920000_daily.parquet") == "920000" + + +# ---------- all good → passed ---------- + +def test_verify_all_good_passes(tmp_path): + """staging 全合法且 fresh → passed=True。""" + staging = tmp_path / "staging" + dates = _recent_dates(3) + # 5 只正常 + 1 只北交所(应被扣分母,不影响通过) + for sym in ("sh600000", "sh600004", "sz000001", "sz300001", "sh688001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + _write_staging(str(staging), "2026", "bj920000_daily.parquet", _good_df(dates)) + + start = dates[0] + result = vi.verify(str(staging), start) + + assert result["passed"] is True + assert result["total"] == 6 + assert result["unsupported_skipped"] == 1 # 北交所 920 + assert result["success"] == 5 + assert result["success_rate"] == 1.0 + assert result["fresh_rate"] == 1.0 + assert result["failed_symbols"] == [] + + +# ---------- fatal cases → failed ---------- + +def test_verify_empty_file_fails(tmp_path): + """空 df(DataValidator 直接判 fatal '数据为空')→ 该 symbol 失败。""" + staging = tmp_path / "staging" + dates = _recent_dates(3) + + # 4 只好 + 1 只空 + for sym in ("sh600000", "sh600004", "sz000001", "sz300001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + _write_staging(str(staging), "2026", "sh688001_daily.parquet", pd.DataFrame( + {"date": [], "open": [], "high": [], "low": [], "close": [], "volume": []} + )) + + result = vi.verify(str(staging), dates[0]) + # 4 好 / 5 总 = 0.8 < 0.95 → fail + assert result["passed"] is False + assert result["success"] == 4 + assert result["total"] == 5 + assert "688001" in result["failed_symbols"] + assert result["success_rate"] < vi.MIN_SUCCESS_RATE + # fatal 样本里有 688001 + assert any(s["symbol"] == "688001" for s in result["fatal_samples"]) + + +def test_verify_zero_price_fails(tmp_path): + """价格<=0(D1 fatal)→ 该 symbol 失败。""" + staging = tmp_path / "staging" + dates = _recent_dates(3) + + for sym in ("sh600000", "sh600004", "sz000001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + # 构造 close<=0 的坏 df + bad = pd.DataFrame({ + "date": pd.to_datetime(dates), + "open": [0.0, 0.0, 0.0], "high": [0.0, 0.0, 0.0], + "low": [0.0, 0.0, 0.0], "close": [0.0, 0.0, 0.0], + "volume": [100, 200, 300], + }) + _write_staging(str(staging), "2026", "sz300001_daily.parquet", bad) + + result = vi.verify(str(staging), dates[0]) + # 3 好 / 4 总 = 0.75 < 0.95 → fail + assert result["passed"] is False + assert "300001" in result["failed_symbols"] + # 样本错误里有 D1 + sample = next(s for s in result["fatal_samples"] if s["symbol"] == "300001") + assert any("D1" in e for e in sample["errors"]) + + +def test_verify_bse_excluded_from_denominator(tmp_path): + """北交所码(920/921/83/87)从分母扣——不算失败也不算成功。""" + staging = tmp_path / "staging" + dates = _recent_dates(3) + # 3 只全北交所 → denom=0 → passed=False(denom=0 算不通过,因为没有有效样本可验) + for sym in ("bj920000", "bj920001", "sz830001", "sz870001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + + result = vi.verify(str(staging), dates[0]) + assert result["unsupported_skipped"] == 4 + assert result["denom"] == 0 + assert result["passed"] is False # denom=0 → 不通过 + + +def test_verify_missing_year_dir_raises(tmp_path): + """staging 的 year 目录不存在 → FileNotFoundError。""" + staging = tmp_path / "staging" + staging.mkdir() + with pytest.raises(FileNotFoundError): + vi.verify(str(staging), "2099-01-01") + + +# ---------- _latest_available_trading_day(收盘感知) ---------- +# 2026-07-13=Mon, 07-14=Tue, 07-10=Fri, 07-11=Sat, 07-12=Sun + +def test_latest_available_premarket_weekday(): + """工作日盘前(<15:00)→ 上一交易日。""" + now = datetime(2026, 7, 14, 4, 12) # 周二 04:12 + assert vi._latest_available_trading_day(now) == "2026-07-13" + + +def test_latest_available_after_market_weekday(): + """工作日盘后(>=15:00)→ 今天。""" + now = datetime(2026, 7, 14, 16, 0) # 周二 16:00 + assert vi._latest_available_trading_day(now) == "2026-07-14" + + +def test_latest_available_at_15_exact(): + """15:00 整点算盘后(>= 15:00 → 今天)。""" + now = datetime(2026, 7, 14, 15, 0) # 周二 15:00 + assert vi._latest_available_trading_day(now) == "2026-07-14" + + +def test_latest_available_weekend(): + """周末 → 上周五。""" + assert vi._latest_available_trading_day(datetime(2026, 7, 11, 10, 0)) == "2026-07-10" # Sat + assert vi._latest_available_trading_day(datetime(2026, 7, 12, 20, 0)) == "2026-07-10" # Sun + + +def test_latest_available_monday_premarket(): + """周一盘前 → 上周五(回退周末)。""" + now = datetime(2026, 7, 13, 4, 0) # 周一 04:00 + assert vi._latest_available_trading_day(now) == "2026-07-10" + + +def test_latest_available_default_now(): + """不传 now → 返回字符串(不抛异常)。""" + result = vi._latest_available_trading_day() + assert isinstance(result, str) + assert len(result) == 10 # YYYY-MM-DD + + +# ---------- verify 收盘感知场景(跨夜 catch-up 核心 bug) ---------- + +def test_verify_premarket_overnight_passes(tmp_path): + """用例1 跨夜/盘前:now=周二 04:12, staging max=周一 07-13 → 目标=07-13 → fresh → passed.""" + staging = tmp_path / "staging" + dates = ["2026-07-09", "2026-07-10", "2026-07-13"] # max=周一 + for sym in ("sh600000", "sh600004", "sz000001", "sz300001", "sh688001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + _write_staging(str(staging), "2026", "bj920000_daily.parquet", _good_df(dates)) + + now = datetime(2026, 7, 14, 4, 12) # 周二 04:12 盘前 + result = vi.verify(str(staging), "2026-07-06", _now=now) + + assert result["latest_trading_day"] == "2026-07-13" + assert result["fresh_rate"] == 1.0 + assert result["passed"] is True + + +def test_verify_after_market_fails_if_stale(tmp_path): + """用例2 盘后:now=周二 16:00, staging max=周一 07-13 → 目标=07-14 → fresh_rate=0 → fail.""" + staging = tmp_path / "staging" + dates = ["2026-07-09", "2026-07-10", "2026-07-13"] # max=周一 + for sym in ("sh600000", "sh600004", "sz000001", "sz300001", "sh688001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + _write_staging(str(staging), "2026", "bj920000_daily.parquet", _good_df(dates)) + + now = datetime(2026, 7, 14, 16, 0) # 周二 16:00 盘后 + result = vi.verify(str(staging), "2026-07-06", _now=now) + + assert result["latest_trading_day"] == "2026-07-14" + assert result["fresh_rate"] == 0.0 + assert result["passed"] is False + + +def test_verify_weekend_passes(tmp_path): + """用例3 周末:now=周六, staging max=周五 07-10 → 目标=07-10 → fresh → passed.""" + staging = tmp_path / "staging" + dates = ["2026-07-08", "2026-07-09", "2026-07-10"] # max=周五 + for sym in ("sh600000", "sh600004", "sz000001", "sz300001", "sh688001"): + _write_staging(str(staging), "2026", f"{sym}_daily.parquet", _good_df(dates)) + _write_staging(str(staging), "2026", "bj920000_daily.parquet", _good_df(dates)) + + now = datetime(2026, 7, 11, 10, 0) # 周六 10:00 + result = vi.verify(str(staging), "2026-07-06", _now=now) + + assert result["latest_trading_day"] == "2026-07-10" + assert result["fresh_rate"] == 1.0 + assert result["passed"] is True