506 lines
18 KiB
Python
506 lines
18 KiB
Python
"""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)
|