Files
claude_dev e3b688354f 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 测试.
2026-07-15 07:12:46 +08:00

180 lines
7.2 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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")