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:
2026-07-15 07:12:46 +08:00
parent 54f9ab4c4f
commit e3b688354f
10 changed files with 1341 additions and 24 deletions
View File
+155
View File
@@ -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 只 fail35%)→ 触发。"""
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 只 fail25%)→ 不触发。"""
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")
+179
View File
@@ -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:
"""构造合法日线 dfdate + 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 新 = 1032 个重复去重后保留 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 只 symbol2 只 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: 每只新增 3staging 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):
"""空 dfDataValidator 直接判 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):
"""价格<=0D1 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=Falsedenom=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