fix(data): data_platform硬化(增量merge/verify+raw_redownload/run_daily_update)+测试
merge_increment/verify_increment 增量staging→验证→合并工具; raw_redownload/run_daily_update/import_vnpy_daily 强化; 补 data_platform 与 index_downloader 测试.
This commit is contained in:
@@ -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 # 每批插入行数
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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) ==="
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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")
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user