From ef09d5b4190ed1d1fccd19852bd5fb352c7f1953 Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 9 Sep 2026 00:32:55 +0800 Subject: [PATCH] =?UTF-8?q?feat(quant):=20M7.3=20=E7=A0=94=E7=A9=B6?= =?UTF-8?q?=E8=A1=8C=E6=83=85=E5=8F=A3=E5=BE=84=E6=98=BE=E5=BC=8F=E5=8C=96?= =?UTF-8?q?=EF=BC=88=E9=BB=98=E8=AE=A4=E4=B8=8D=E5=A4=8D=E6=9D=83=20none?= =?UTF-8?q?=EF=BC=8C=E5=8F=AF=E5=88=87=20qfq=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 通过 --- .../application/services/selection_service.py | 1 + backend/app/domain/entities/research.py | 4 + backend/app/domain/entities/selection.py | 4 + backend/app/domain/repositories/market.py | 11 +- .../sqlalchemy/repositories/market_impl.py | 16 ++- backend/app/quant/service.py | 14 ++- backend/tests/test_api.py | 2 +- backend/tests/test_jobs_experiments.py | 2 +- backend/tests/test_price_adjustment.py | 106 ++++++++++++++++++ backend/tests/test_selection.py | 2 +- .../test_selection_backtest_consistency.py | 2 +- backend/tests/test_selection_condition.py | 2 +- 12 files changed, 153 insertions(+), 13 deletions(-) create mode 100644 backend/tests/test_price_adjustment.py diff --git a/backend/app/application/services/selection_service.py b/backend/app/application/services/selection_service.py index ee8d84b..0478c32 100644 --- a/backend/app/application/services/selection_service.py +++ b/backend/app/application/services/selection_service.py @@ -62,6 +62,7 @@ class SelectionService: as_of - timedelta(days=query.warmup_days), as_of, columns, + adjust=query.price_adjustment, ) financial: dict[str, FinancialIndicator] = {} if query.method == "condition" and self._uses_fundamental(query): diff --git a/backend/app/domain/entities/research.py b/backend/app/domain/entities/research.py index 3f9a95e..21e48a3 100644 --- a/backend/app/domain/entities/research.py +++ b/backend/app/domain/entities/research.py @@ -62,6 +62,10 @@ class ResearchSpec(BaseModel): type: str = Field(default="backtest", pattern="^(factor_test|backtest)$") universe: UniverseSpec = UniverseSpec() + price_adjustment: str = Field( + default="none", pattern="^(none|qfq)$", + description="研究行情口径:none 不复权(默认)/ qfq 前复权(result 与 config_snapshot 中显式)", + ) factors: list[FactorSpec] = Field(min_length=1) selection: SelectionSpec = SelectionSpec() rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$") diff --git a/backend/app/domain/entities/selection.py b/backend/app/domain/entities/selection.py index 2c3ffc3..8bf6d46 100644 --- a/backend/app/domain/entities/selection.py +++ b/backend/app/domain/entities/selection.py @@ -25,6 +25,10 @@ class SelectionQuery(BaseModel): """一次选股查询(v2 §14.2 Selection 输入)。""" universe: UniverseSpec = UniverseSpec() + price_adjustment: str = Field( + default="none", pattern="^(none|qfq)$", + description="行情口径:none 不复权(默认)/ qfq 前复权(结果 config_snapshot 中显式)", + ) # 研究时点:None → 引擎用 <= 今天最近可用交易日;显式给历史日期即做历史选股 as_of: date | None = Field( default=None, description="选股时点;历史回测/解释用具体日期,当前选股可留空" diff --git a/backend/app/domain/repositories/market.py b/backend/app/domain/repositories/market.py index f9c2e94..a2340fe 100644 --- a/backend/app/domain/repositories/market.py +++ b/backend/app/domain/repositories/market.py @@ -42,8 +42,14 @@ class DailyBarRepository(Protocol): def get_range(self, symbol: str, start: date, end: date) -> list[DailyBar]: ... - def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]: - """批量区间查询(研究服务装配面板用,避免逐只查询)。""" + def get_range_many( + self, + symbols: Sequence[str], + start: date, + end: date, + adjust: str = "none", + ) -> list[DailyBar]: + """批量区间查询(研究装配面板用);adjust 指定行情口径(none 不复权 / qfq)。""" def stream_range_many_columns( self, @@ -51,6 +57,7 @@ class DailyBarRepository(Protocol): start: date, end: date, columns: Sequence[str], + adjust: str = "none", ) -> Iterator[tuple]: """流式(分批 yield)返回 symbol, trade_date(iso str), 数值列(float) 元组。 diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py index a417982..156b8d2 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py @@ -158,16 +158,24 @@ class SqlAlchemyDailyBarRepository: ).all() return [DailyBar.model_validate(r, from_attributes=True) for r in rows] - def get_range_many(self, symbols: Sequence[str], start: date, end: date) -> list[DailyBar]: - rows = self._session.scalars( + def get_range_many( + self, + symbols: Sequence[str], + start: date, + end: date, + adjust: str = "none", + ) -> list[DailyBar]: + stmt = ( select(StockDailyModel) .where( StockDailyModel.symbol.in_(list(symbols)), StockDailyModel.trade_date >= start, StockDailyModel.trade_date <= end, + StockDailyModel.adjust == adjust, # 研究主口径:不复权(v2 §8) ) .order_by(StockDailyModel.trade_date) - ).all() + ) + rows = self._session.scalars(stmt).all() return [DailyBar.model_validate(r, from_attributes=True) for r in rows] def stream_range_many_columns( @@ -176,6 +184,7 @@ class SqlAlchemyDailyBarRepository: start: date, end: date, columns: Sequence[str], + adjust: str = "none", ) -> Iterator[tuple]: """流式返回 (symbol, trade_date_iso, *float_cols) 元组,分批拉取。 @@ -193,6 +202,7 @@ class SqlAlchemyDailyBarRepository: StockDailyModel.symbol.in_(list(symbols)), StockDailyModel.trade_date >= start, StockDailyModel.trade_date <= end, + StockDailyModel.adjust == adjust, ) .order_by(StockDailyModel.symbol, StockDailyModel.trade_date) .execution_options(yield_per=20000) diff --git a/backend/app/quant/service.py b/backend/app/quant/service.py index bc82479..40a0964 100644 --- a/backend/app/quant/service.py +++ b/backend/app/quant/service.py @@ -75,6 +75,7 @@ def load_daily_df( start: date, end: date, columns: list[str], + adjust: str = "none", ) -> pd.DataFrame: """从 Repository 装配行情长表(供研究/选股共用)。 @@ -86,12 +87,14 @@ def load_daily_df( streamer = getattr(daily_repo, "stream_range_many_columns", None) if streamer is not None: try: - return _frame_from_stream(streamer(symbols, start, end, sorted(columns)), sorted(columns)) + return _frame_from_stream( + streamer(symbols, start, end, sorted(columns), adjust=adjust), sorted(columns) + ) except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现) pass get_many = getattr(daily_repo, "get_range_many", None) if get_many is not None: - bars = list(get_many(symbols, start, end)) + bars = list(get_many(symbols, start, end, adjust=adjust)) else: # 兜底:逐只查询 bars = [] for sym in symbols: @@ -134,5 +137,10 @@ class ResearchService: # 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV) required = self._engine.required_columns(spec) return load_daily_df( - self._daily_repo, [s.symbol for s in stocks], data_start, end, sorted(required) + self._daily_repo, + [s.symbol for s in stocks], + data_start, + end, + sorted(required), + adjust=spec.price_adjustment, ) diff --git a/backend/tests/test_api.py b/backend/tests/test_api.py index b02c8fa..48bfc49 100644 --- a/backend/tests/test_api.py +++ b/backend/tests/test_api.py @@ -48,7 +48,7 @@ def client(tmp_path) -> TestClient: bars = bars_dataframe_to_daily_bars(daily_df) class _MemDailyRepo: - def get_range_many(self, symbols, start, end): + def get_range_many(self, symbols, start, end, adjust="none"): out = [] for b in bars: if b.symbol in symbols and start <= b.trade_date <= end: diff --git a/backend/tests/test_jobs_experiments.py b/backend/tests/test_jobs_experiments.py index cb3da60..840a071 100644 --- a/backend/tests/test_jobs_experiments.py +++ b/backend/tests/test_jobs_experiments.py @@ -59,7 +59,7 @@ class _FakeDailyRepo: def __init__(self, bars): self._bars = bars - def get_range_many(self, symbols, start, end): + def get_range_many(self, symbols, start, end, adjust="none"): return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end] def get_range(self, symbol, start, end): diff --git a/backend/tests/test_price_adjustment.py b/backend/tests/test_price_adjustment.py new file mode 100644 index 0000000..ec889de --- /dev/null +++ b/backend/tests/test_price_adjustment.py @@ -0,0 +1,106 @@ +"""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" diff --git a/backend/tests/test_selection.py b/backend/tests/test_selection.py index a3fc283..572abce 100644 --- a/backend/tests/test_selection.py +++ b/backend/tests/test_selection.py @@ -69,7 +69,7 @@ class _MemDailyRepo: def get_range(self, symbol, start, end) -> list[DailyBar]: return self._bars([symbol], start, end) - def get_range_many(self, symbols, start, end) -> list[DailyBar]: + def get_range_many(self, symbols, start, end, adjust="none") -> list[DailyBar]: return self._bars(list(symbols), start, end) def latest_date(self, symbol: str) -> date | None: diff --git a/backend/tests/test_selection_backtest_consistency.py b/backend/tests/test_selection_backtest_consistency.py index a8574ce..60d3ccf 100644 --- a/backend/tests/test_selection_backtest_consistency.py +++ b/backend/tests/test_selection_backtest_consistency.py @@ -42,7 +42,7 @@ class _MemDailyRepo: def get_range(self, symbol, start, end): return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end] - def get_range_many(self, symbols, start, end): + def get_range_many(self, symbols, start, end, adjust="none"): syms = set(symbols) return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end] diff --git a/backend/tests/test_selection_condition.py b/backend/tests/test_selection_condition.py index 7a0ba14..be97d4b 100644 --- a/backend/tests/test_selection_condition.py +++ b/backend/tests/test_selection_condition.py @@ -38,7 +38,7 @@ class _MemDailyRepo: def get_range(self, symbol, start, end): return [b for b in self._bars_all if b.symbol == symbol and start <= b.trade_date <= end] - def get_range_many(self, symbols, start, end): + def get_range_many(self, symbols, start, end, adjust="none"): syms = set(symbols) return [b for b in self._bars_all if b.symbol in syms and start <= b.trade_date <= end]