Files
quant_engine/tests/test_data_adapter.py
T
George BerkshireandMavis c04acf0ab6 feat: quant_engine 独立量化引擎(v1.2.0 重构 bootstrap)
从 research_results 抽出纯回测核心能力,形成独立成员仓。

## 模块(10 个核心)

| 模块 | 内容 |
|---|---|
| alpha_factors | 158 alpha 公式 + 24 基础算子(移植自 qlib alpha158)+ JSONB 工具 |
| execution | 执行仿真(成本/滑点/T+1/涨跌停/部分成交/价差)+ 多日 NAV + PnL 拆解 |
| indicators | 50+ 技术指标(MACD/KDJ/布林/ATR/ADX/等) |
| data_adapter | 桥接 qtdb_pro 长表与新模块(rename/long-wide/复权/vwap 代理) |
| backtest | weight-based 多日仿真 |
| metrics / perf_stats | 绩效指标 |
| factor_library | 通用方法(turnover/IC/winsorize/OLS) |
| portfolio_decomp / risk | 组合分解 + 风险指标 |
| logging | 统一 logger(标准库 + 可选 loguru) |

## 设计原则

- 零重型依赖(numpy/pandas/scipy)
- mypy strict 0 errors(12 source files)
- 394 tests passed(从 research_results 复制 + 适配)
- ruff clean

## 与 research_results 的关系

- research_results 通过 re-export wrapper 保持向后兼容(src.shared.X → quant_engine.X)
- 47 proj 的 import 路径暂时不变,后续逐步迁移
- 本仓角色:researchhub_workspace 引擎层

Co-Authored-By: Mavis <noreply@mavis.local>
2026-08-19 15:50:01 +08:00

322 lines
12 KiB
Python
Raw 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 (
add_vwap_proxy,
apply_adj_factor,
load_qtdb_daily,
long_to_wide,
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_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)
# ── 端到端:长表 → 适配 → 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)