e3b688354f
merge_increment/verify_increment 增量staging→验证→合并工具; raw_redownload/run_daily_update/import_vnpy_daily 强化; 补 data_platform 与 index_downloader 测试.
180 lines
7.2 KiB
Python
180 lines
7.2 KiB
Python
"""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")
|