- 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 通过
107 lines
4.1 KiB
Python
107 lines
4.1 KiB
Python
"""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"
|