Files
quant_engine/tests/test_data_adapter.py
2026-08-26 20:55:09 +08:00

506 lines
18 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.
"""src/shared/data_adapter.py 单元测试(v1.2.0 数据对齐层)。"""
from __future__ import annotations
import numpy as np
import pandas as pd
import pytest
from quant_engine.data_adapter import (
AssetReturnSnapshot,
add_vwap_proxy,
apply_adj_factor,
load_qtdb_daily,
long_to_wide,
prepare_asset_return_snapshot,
prepare_execution_inputs,
prepare_stock_series,
rename_tushare_columns,
wide_to_long,
)
@pytest.fixture
def tushare_long() -> pd.DataFrame:
"""模拟 qtdb_pro.hq_daily 长表(Tushare 原始命名)。"""
return pd.DataFrame(
{
"ts_code": ["000001.SZ", "000001.SZ", "600000.SH", "600000.SH"],
"trade_date": ["2024-01-01", "2024-01-02", "2024-01-01", "2024-01-02"],
"open": [10.0, 11.0, 20.0, 21.0],
"high": [10.5, 11.5, 20.5, 21.5],
"low": [9.8, 10.8, 19.8, 20.8],
"close": [10.2, 11.2, 20.2, 21.2],
"pre_close": [10.0, 10.2, 20.0, 20.2],
"pct_chg": [2.0, 9.8, 1.0, 4.95],
"vol": [1000.0, 1100.0, 2000.0, 2100.0],
"amount": [10000.0, 12000.0, 40000.0, 44000.0],
}
)
# ── long_to_wide / wide_to_long ──────────────────────────────
def test_long_to_wide_basic(tushare_long: pd.DataFrame) -> None:
"""长表 → 宽表:2 股票 × 2 日期。"""
renamed = rename_tushare_columns(tushare_long)
wide = long_to_wide(renamed, value_col="close")
assert wide.shape == (2, 2)
assert set(wide.columns) == {"000001.SZ", "600000.SH"}
assert wide.index.is_monotonic_increasing
# 值校验:000001.SZ 两天 close
assert wide.loc[wide.index[0], "000001.SZ"] == pytest.approx(10.2)
assert wide.loc[wide.index[1], "000001.SZ"] == pytest.approx(11.2)
def test_long_to_wide_empty() -> None:
"""空输入 → 空 DataFrame。"""
assert long_to_wide(pd.DataFrame()).empty
def test_long_to_wide_missing_col_raises(tushare_long: pd.DataFrame) -> None:
"""缺列应报错。"""
with pytest.raises(ValueError, match="缺少列"):
long_to_wide(tushare_long, value_col="not_exist")
def test_wide_to_long_roundtrip(tushare_long: pd.DataFrame) -> None:
"""wide → long → wide roundtrip 应一致。"""
renamed = rename_tushare_columns(tushare_long)
wide = long_to_wide(renamed, value_col="close")
long_back = wide_to_long(wide, value_name="close")
# 重建 wide 应与原一致(NaN 会被 dropna)
wide2 = long_to_wide(long_back, value_col="close")
pd.testing.assert_frame_equal(
wide,
wide2,
check_names=False,
)
def test_wide_to_long_columns(tushare_long: pd.DataFrame) -> None:
"""wide_to_long 输出三列。"""
renamed = rename_tushare_columns(tushare_long)
wide = long_to_wide(renamed, value_col="close")
long_df = wide_to_long(wide, value_name="close")
assert set(long_df.columns) == {"trade_date", "stock_code", "close"}
# ── rename_tushare_columns ──────────────────────────────
def test_rename_tushare_columns_basic(tushare_long: pd.DataFrame) -> None:
"""ts_code→stock_code, vol→volume。"""
renamed = rename_tushare_columns(tushare_long)
assert "stock_code" in renamed.columns
assert "volume" in renamed.columns
assert "ts_code" not in renamed.columns
assert "vol" not in renamed.columns
assert "close" in renamed.columns # 已一致,不变
def test_rename_tushare_columns_noop() -> None:
"""无 Tushare 列 → 原样。"""
df = pd.DataFrame({"a": [1], "b": [2]})
assert rename_tushare_columns(df).columns.tolist() == ["a", "b"]
# ── add_vwap_proxy ──────────────────────────────────────
def test_vwap_amount_vol(tushare_long: pd.DataFrame) -> None:
"""vwap = amount*10/volume(amount 千元,vol 手)。"""
renamed = rename_tushare_columns(tushare_long)
out = add_vwap_proxy(renamed, method="amount_vol")
assert "vwap" in out.columns
# 000001.SZ 第一天:amount=10000(千元), volume=1000(手)
# vwap = 10000 * 10 / 1000 = 100 元/股
first = out[out["stock_code"] == "000001.SZ"].iloc[0]
assert first["vwap"] == pytest.approx(100.0)
def test_vwap_typical(tushare_long: pd.DataFrame) -> None:
"""typical price 代理。"""
renamed = rename_tushare_columns(tushare_long)
out = add_vwap_proxy(renamed, method="typical")
first = out[out["stock_code"] == "000001.SZ"].iloc[0]
expected = (10.5 + 9.8 + 10.2) / 3
assert first["vwap"] == pytest.approx(expected)
def test_vwap_close_fallback(tushare_long: pd.DataFrame) -> None:
"""close 兜底。"""
renamed = rename_tushare_columns(tushare_long)
out = add_vwap_proxy(renamed, method="close")
first = out[out["stock_code"] == "000001.SZ"].iloc[0]
assert first["vwap"] == pytest.approx(10.2)
def test_vwap_insufficient_cols_raises() -> None:
"""列不足应报错。"""
df = pd.DataFrame({"a": [1]})
with pytest.raises(ValueError, match="列不足"):
add_vwap_proxy(df)
# ── apply_adj_factor ──────────────────────────────────────
def test_apply_adj_factor_qfq() -> None:
"""前复权:价格 × adj / 最新 adj。"""
daily = pd.DataFrame(
{
"stock_code": ["A", "A", "A"],
"trade_date": ["2024-01-01", "2024-01-02", "2024-01-03"],
"close": [10.0, 20.0, 20.0], # 1/2 除权(10→20 拆股?反向演示)
"open": [10.0, 20.0, 20.0],
"high": [11.0, 21.0, 21.0],
"low": [9.0, 19.0, 19.0],
}
)
adj = pd.DataFrame(
{
"stock_code": ["A", "A", "A"],
"trade_date": ["2024-01-01", "2024-01-02", "2024-01-03"],
"adj_factor": [2.0, 2.0, 1.0], # 最新 = 1.0
}
)
out = apply_adj_factor(daily, adj, mode="qfq")
# 1/1: 10 * 2/1 = 20;1/2: 20 * 2/1 = 40;1/3: 20 * 1/1 = 20
closes = out.set_index("trade_date")["close"]
assert closes["2024-01-01"] == pytest.approx(20.0)
assert closes["2024-01-02"] == pytest.approx(40.0)
assert closes["2024-01-03"] == pytest.approx(20.0)
def test_apply_adj_factor_hfq() -> None:
"""后复权:价格 × adj。"""
daily = pd.DataFrame(
{
"stock_code": ["A", "A"],
"trade_date": ["2024-01-01", "2024-01-02"],
"close": [10.0, 10.0],
}
)
adj = pd.DataFrame(
{
"stock_code": ["A", "A"],
"trade_date": ["2024-01-01", "2024-01-02"],
"adj_factor": [2.0, 4.0],
}
)
out = apply_adj_factor(daily, adj, mode="hfq")
closes = out.set_index("trade_date")["close"]
assert closes["2024-01-01"] == pytest.approx(20.0)
assert closes["2024-01-02"] == pytest.approx(40.0)
def test_apply_adj_factor_empty() -> None:
"""adj_factor 为空 → 返回原行情。"""
daily = pd.DataFrame({"stock_code": ["A"], "trade_date": ["2024-01-01"], "close": [10.0]})
out = apply_adj_factor(daily, pd.DataFrame())
assert out["close"].iloc[0] == pytest.approx(10.0)
def test_apply_adj_factor_missing_dates() -> None:
"""无因子日期 → 保持原值。"""
daily = pd.DataFrame(
{
"stock_code": ["A", "A"],
"trade_date": ["2024-01-01", "2024-01-02"],
"close": [10.0, 12.0],
}
)
adj = pd.DataFrame(
{
"stock_code": ["A"],
"trade_date": ["2024-01-01"],
"adj_factor": [1.0],
}
)
out = apply_adj_factor(daily, adj, mode="qfq")
closes = out.set_index("trade_date")["close"]
assert closes["2024-01-01"] == pytest.approx(10.0)
assert closes["2024-01-02"] == pytest.approx(12.0) # 无因子 → 原值
# ── prepare_stock_series ──────────────────────────────────
def test_prepare_stock_series_basic(tushare_long: pd.DataFrame) -> None:
"""单股提取 → Series dict。"""
renamed = rename_tushare_columns(tushare_long)
series_map = prepare_stock_series(renamed, "000001.SZ")
assert "close" in series_map
assert "open" in series_map
assert "high" in series_map
assert "low" in series_map
assert "volume" in series_map
assert "vwap" in series_map # 自动补
assert len(series_map["close"]) == 2
assert series_map["close"].iloc[0] == pytest.approx(10.2)
def test_prepare_stock_series_unknown_stock(tushare_long: pd.DataFrame) -> None:
"""未知股票 → 空 dict。"""
renamed = rename_tushare_columns(tushare_long)
assert prepare_stock_series(renamed, "999999.SZ") == {}
def test_prepare_stock_series_empty() -> None:
"""空输入 → 空 dict。"""
assert prepare_stock_series(pd.DataFrame(), "A") == {}
# ── prepare_execution_inputs ──────────────────────────────
def test_prepare_execution_inputs_basic(tushare_long: pd.DataFrame) -> None:
"""prices + volumes 宽表。"""
renamed = rename_tushare_columns(tushare_long)
prices, volumes = prepare_execution_inputs(renamed)
assert prices.shape == (2, 2)
assert volumes.shape == (2, 2)
# prices 值 = close
assert prices.iloc[0, 0] == pytest.approx(10.2)
# volumes 值 = volume
assert volumes.iloc[0, 0] == pytest.approx(1000.0)
def test_prepare_execution_inputs_can_select_next_session_open_price(
tushare_long: pd.DataFrame,
) -> None:
"""显式 price_col=open 时应生成开盘执行价矩阵。"""
renamed = rename_tushare_columns(tushare_long)
prices, _volumes = prepare_execution_inputs(renamed, price_col="open")
assert prices.iloc[0, 0] == pytest.approx(10.0)
def test_prepare_execution_inputs_no_volume() -> None:
"""无 volume 列 → volumes 全 1.0。"""
df = pd.DataFrame(
{
"stock_code": ["A"],
"trade_date": ["2024-01-01"],
"close": [10.0],
}
)
_prices, volumes = prepare_execution_inputs(df)
assert volumes.iloc[0, 0] == pytest.approx(1.0)
def test_prepare_execution_inputs_missing_close_raises() -> None:
"""缺 close 列应报错。"""
df = pd.DataFrame({"stock_code": ["A"], "trade_date": ["2024-01-01"]})
with pytest.raises(ValueError, match="缺 close"):
prepare_execution_inputs(df)
def test_prepare_execution_inputs_missing_selected_price_raises() -> None:
df = pd.DataFrame({"stock_code": ["A"], "trade_date": ["2024-01-01"], "close": [10.0]})
with pytest.raises(ValueError, match="缺 open"):
prepare_execution_inputs(df, price_col="open")
# ── prepare_asset_return_snapshot ────────────────────────────
def _daily_prices() -> pd.DataFrame:
return pd.DataFrame(
{
"stock_code": ["B", "A", "B", "A", "B", "A"],
"trade_date": [
"2024-01-02",
"2024-01-01",
"2024-01-01",
"2024-01-03",
"2024-01-03",
"2024-01-02",
],
"close": [18.0, 10.0, 20.0, 12.1, 19.8, 11.0],
}
)
def test_prepare_asset_return_snapshot_is_stable_and_immutable_by_interface() -> None:
snapshot = prepare_asset_return_snapshot(
_daily_prices(),
source="qtdb_pro.hq_daily",
source_snapshot_id="hq-daily:2024-01-03:v1",
adjustment="qfq",
)
assert isinstance(snapshot, AssetReturnSnapshot)
assert snapshot.data_snapshot_id.startswith("asset-returns-v1:")
assert snapshot.source == "qtdb_pro.hq_daily"
assert snapshot.source_snapshot_id == "hq-daily:2024-01-03:v1"
assert snapshot.price_field == "close"
assert snapshot.adjustment == "qfq"
assert snapshot.return_method == "simple"
assert snapshot.start_date.isoformat() == "2024-01-01"
assert snapshot.end_date.isoformat() == "2024-01-03"
assert snapshot.sessions == 3
assert snapshot.assets == ("A", "B")
expected = pd.DataFrame(
{
"A": [np.nan, 0.1, 0.1],
"B": [np.nan, -0.1, 0.1],
},
index=pd.to_datetime(["2024-01-01", "2024-01-02", "2024-01-03"]),
)
expected.index.name = "trade_date"
expected.columns.name = "stock_code"
pd.testing.assert_frame_equal(snapshot.returns, expected)
exposed = snapshot.returns
exposed.iloc[1, 0] = 999.0
assert snapshot.returns.iloc[1, 0] == pytest.approx(0.1)
def test_asset_return_snapshot_identity_is_order_independent_and_content_addressed() -> None:
kwargs = {
"source": "qtdb_pro.hq_daily",
"source_snapshot_id": "hq-daily:2024-01-03:v1",
"adjustment": "none",
}
baseline = prepare_asset_return_snapshot(_daily_prices(), **kwargs)
shuffled = prepare_asset_return_snapshot(
_daily_prices().sample(frac=1.0, random_state=7),
**kwargs,
)
changed_prices = _daily_prices().copy()
changed_prices.loc[changed_prices["close"] == 12.1, "close"] = 12.2
changed_content = prepare_asset_return_snapshot(changed_prices, **kwargs)
changed_source = prepare_asset_return_snapshot(
_daily_prices(),
source="qtdb_pro.hq_daily",
source_snapshot_id="hq-daily:2024-01-03:v2",
adjustment="none",
)
assert shuffled.data_snapshot_id == baseline.data_snapshot_id
assert changed_content.data_snapshot_id != baseline.data_snapshot_id
assert changed_source.data_snapshot_id != baseline.data_snapshot_id
def test_prepare_asset_return_snapshot_does_not_fill_missing_prices() -> None:
prices = _daily_prices()
prices.loc[
(prices["stock_code"] == "A") & (prices["trade_date"] == "2024-01-02"),
"close",
] = np.nan
snapshot = prepare_asset_return_snapshot(
prices,
source="qtdb_pro.hq_daily",
source_snapshot_id="hq-daily:missing-middle",
)
assert pd.isna(snapshot.returns.loc[pd.Timestamp("2024-01-02"), "A"])
assert pd.isna(snapshot.returns.loc[pd.Timestamp("2024-01-03"), "A"])
def test_prepare_asset_return_snapshot_rejects_duplicate_sessions() -> None:
duplicate = pd.concat([_daily_prices(), _daily_prices().iloc[[0]]], ignore_index=True)
with pytest.raises(ValueError, match="duplicate"):
prepare_asset_return_snapshot(
duplicate,
source="qtdb_pro.hq_daily",
source_snapshot_id="hq-daily:duplicate",
)
@pytest.mark.parametrize("invalid_price", [0.0, -1.0, np.inf])
def test_prepare_asset_return_snapshot_rejects_invalid_prices(invalid_price: float) -> None:
prices = _daily_prices()
prices.loc[0, "close"] = invalid_price
with pytest.raises(ValueError, match="positive finite"):
prepare_asset_return_snapshot(
prices,
source="qtdb_pro.hq_daily",
source_snapshot_id="hq-daily:invalid-price",
)
@pytest.mark.parametrize(
("source", "source_snapshot_id", "adjustment"),
[
("", "source-1", "none"),
("qtdb_pro.hq_daily", "", "none"),
("qtdb_pro.hq_daily", "source-1", ""),
],
)
def test_prepare_asset_return_snapshot_requires_explicit_identity_semantics(
source: str,
source_snapshot_id: str,
adjustment: str,
) -> None:
with pytest.raises(ValueError, match="must be non-empty"):
prepare_asset_return_snapshot(
_daily_prices(),
source=source,
source_snapshot_id=source_snapshot_id,
adjustment=adjustment,
)
def test_asset_return_snapshot_feeds_reproducible_covariance_lineage() -> None:
from quant_engine.risk import estimate_covariance_snapshot
market_snapshot = prepare_asset_return_snapshot(
_daily_prices(),
source="qtdb_pro.hq_daily",
source_snapshot_id="hq-daily:2024-01-03:v1",
)
covariance_snapshot = estimate_covariance_snapshot(
market_snapshot.returns,
as_of_date=market_snapshot.end_date,
lookback_sessions=3,
min_observations=2,
data_snapshot_id=market_snapshot.data_snapshot_id,
)
assert covariance_snapshot.data_snapshot_id == market_snapshot.data_snapshot_id
assert covariance_snapshot.snapshot_id.startswith("sample-cov-v1:")
# ── 端到端:长表 → 适配 → alpha158 + execution ──────────────
def test_end_to_end_tushare_to_alpha158(tushare_long: pd.DataFrame) -> None:
"""真实链路:Tushare 长表 → 单股 Series → alpha158。"""
from quant_engine.alpha_factors import alpha_001, alpha_005
renamed = rename_tushare_columns(tushare_long)
series_map = prepare_stock_series(renamed, "000001.SZ")
# alpha_001: rank(ts_rank(close, 5))
r1 = alpha_001(series_map["close"])
assert isinstance(r1, pd.Series)
# alpha_005: correlation(close, volume, 10)
r5 = alpha_005(series_map["close"], series_map["volume"])
assert isinstance(r5, pd.Series)
def test_end_to_end_tushare_to_execution(tushare_long: pd.DataFrame) -> None:
"""真实链路:Tushare 长表 → execution 宽表 → 端到端 POC。"""
from quant_engine.execution import simulate_with_daily_data
renamed = rename_tushare_columns(tushare_long)
prices, _volumes = prepare_execution_inputs(renamed)
# 用较长的数据才能跑(2 天太短,先验证不崩 + 返回正确类型)
positions = simulate_with_daily_data(prices, initial_cash=1_000_000.0)
assert len(positions) == len(prices)
def test_load_qtdb_daily_offline() -> None:
"""无真实 ClickHouse 时返回空表(不崩)。"""
df = load_qtdb_daily(["000001.SZ"], "2024-01-01")
# 本环境无 .env → 连接失败 → 空表
assert isinstance(df, pd.DataFrame)