"""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")