This commit was merged in pull request #7.
This commit is contained in:
@@ -7,10 +7,12 @@ 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,
|
||||
@@ -266,6 +268,17 @@ def test_prepare_execution_inputs_basic(tushare_long: pd.DataFrame) -> None:
|
||||
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(
|
||||
@@ -286,6 +299,177 @@ def test_prepare_execution_inputs_missing_close_raises() -> None:
|
||||
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 ──────────────
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user