Files
qlib/backend/tests/test_price_adjustment.py
T
Simon ef09d5b419 feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数
  (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险
- ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot
  落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数
- tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、
  spec 默认与字段记录;全量 pytest 通过
2026-09-09 00:32:55 +08:00

107 lines
4.1 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.
"""M7.3 行情口径测试:Repository 读路径按 adjust 过滤(不复权主口径),spec 记录口径。
混合行场景:同 symbol/date 存在 tushare/none 与 sina/qfq 行时,研究读取
(get_range_many / stream)默认只取 adjust=none —— 消除「混合口径污染因子」风险。
"""
from __future__ import annotations
from datetime import date
from decimal import Decimal
from app.domain.entities.market import DailyBar
from app.domain.entities.research import ResearchSpec
from app.domain.entities.selection import SelectionQuery
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
)
from sqlalchemy import create_engine
from sqlalchemy.orm import Session
_D = date(2024, 6, 3)
_D2 = date(2024, 6, 4)
_D3 = date(2024, 6, 5)
def _bar(adjust: str, close: str, source: str = "tushare", day=None) -> DailyBar:
return DailyBar(
symbol="600519.SH",
trade_date=day or _D,
source=source,
adjust=adjust,
open=Decimal("100"), high=Decimal("101"), low=Decimal("99"),
close=Decimal(close), volume=Decimal("1000"), amount=Decimal("100000"),
)
class TestAdjustFilter:
def _session(self, tmp_path) -> Session:
engine = create_engine(f"sqlite:///{tmp_path / 'adj.db'}", future=True)
Base.metadata.create_all(engine)
return Session(engine)
def test_get_range_many_filters_adjust(self, tmp_path) -> None:
with self._session(tmp_path) as session:
repo = SqlAlchemyDailyBarRepository(session)
# 唯一键 (symbol, trade_date):同键共存不可能 —— 用连续三天模拟
# none 主口径两天 + sina/qfq 兜底一天
repo.upsert_many(
[
_bar("none", "1700", day=_D),
_bar("none", "1710", day=_D2),
_bar("qfq", "1680", source="sina", day=_D3),
]
)
session.commit()
none_rows = repo.get_range_many(["600519.SH"], _D, _D3, adjust="none")
assert len(none_rows) == 2 and {float(r.close) for r in none_rows} == {1700, 1710}
qfq_rows = repo.get_range_many(["600519.SH"], _D, _D3, adjust="qfq")
assert len(qfq_rows) == 1 and float(qfq_rows[0].close) == 1680
def test_stream_filters_adjust(self, tmp_path) -> None:
with self._session(tmp_path) as session:
repo = SqlAlchemyDailyBarRepository(session)
repo.upsert_many(
[
_bar("none", "1700", day=_D),
_bar("none", "1710", day=_D2),
_bar("qfq", "1680", source="sina", day=_D3),
]
)
session.commit()
rows = list(
repo.stream_range_many_columns(
["600519.SH"], _D, _D3, ["close"], adjust="none"
)
)
assert len(rows) == 2 and {float(r[-1]) for r in rows} == {1700, 1710}
qrows = list(
repo.stream_range_many_columns(
["600519.SH"], _D, _D3, ["close"], adjust="qfq"
)
)
assert len(qrows) == 1 and float(qrows[0][-1]) == 1680
class TestSpecRecordsAdjustment:
def test_research_spec_default_and_field(self) -> None:
spec = ResearchSpec(
type="backtest",
factors=[{"name": "momentum_60", "weight": 1}],
period=(date(2024, 1, 1), date(2024, 6, 1)),
)
assert spec.price_adjustment == "none"
snap = spec.model_dump(mode="json")
assert snap["price_adjustment"] == "none" # 结果 config_snapshot 可溯源
def test_selection_query_adjustment_in_snapshot(self) -> None:
q = SelectionQuery(factors=[{"name": "momentum_60", "weight": 1}], top_n=5)
assert q.price_adjustment == "none"
assert q.model_dump()["price_adjustment"] == "none"
q2 = SelectionQuery(
factors=[{"name": "momentum_60", "weight": 1}], top_n=5, price_adjustment="qfq"
)
assert q2.price_adjustment == "qfq"