diff --git a/backend/app/domain/entities/research.py b/backend/app/domain/entities/research.py index bf531e4..3f9a95e 100644 --- a/backend/app/domain/entities/research.py +++ b/backend/app/domain/entities/research.py @@ -16,12 +16,20 @@ from pydantic import BaseModel, Field, field_validator, model_validator class UniverseSpec(BaseModel): - """股票池口径。MVP:市场 + 过滤条件;指数成分等 Phase 3 扩展。""" + """股票池口径。MVP:市场 + 过滤条件;指数成分等 Phase 3 扩展。 - market: str = Field(default="CN_A", description="CN_A / CN_B / ...") + symbols 白名单:非空时仅这些股票参与(再叠加其余过滤);供自选池/测试使用。 + market 目前为预留字段(stock.market 存储主板/创业板/科创板等中文枚举,过滤未启用)。 + """ + + market: str = Field(default="CN_A", description="CN_A / CN_B / ...(预留)") exclude_st: bool = True exclude_suspended: bool = True min_listing_days: int = Field(default=250, ge=0, description="上市至少 N 个自然日") + symbols: list[str] = Field( + default_factory=list, + description="白名单(可选):非空时仅这些 symbol 参与选股/回测", + ) class FactorSpec(BaseModel): diff --git a/backend/app/quant/service.py b/backend/app/quant/service.py index 738b0da..bc82479 100644 --- a/backend/app/quant/service.py +++ b/backend/app/quant/service.py @@ -15,41 +15,22 @@ from datetime import date, timedelta import pandas as pd -from app.domain.entities.market import Stock from app.domain.entities.research import ( BacktestResult, FactorTestReport, ResearchSpec, - UniverseSpec, ) from app.domain.repositories.market import ( DailyBarRepository, StockRepository, ) from app.quant.engine import QuantEngine +from app.quant.universe import filter_stocks # noqa: F401 —— 选股/回测共用范围过滤 # 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值) _FRAME_CHUNK_ROWS = 50_000 -def filter_stocks(stocks: list[Stock], universe: UniverseSpec, as_of: date) -> list[Stock]: - """按股票池口径过滤(名称含 ST 判定 —— 名称快照为当日口径,属历史可追溯数据)。""" - out: list[Stock] = [] - for s in stocks: - if s.delist_date is not None and s.delist_date < as_of: - continue - if universe.exclude_st and s.name and "ST" in s.name.upper(): - continue - if ( - universe.min_listing_days - and s.list_date - and (as_of - s.list_date).days < universe.min_listing_days - ): - continue - out.append(s) - return out - - def bars_to_daily_df(bars) -> pd.DataFrame: """DailyBar 列表 → 引擎长表 DataFrame(symbol/trade_date/ohlc/volume/amount)。 diff --git a/backend/app/quant/universe.py b/backend/app/quant/universe.py new file mode 100644 index 0000000..9a31ca5 --- /dev/null +++ b/backend/app/quant/universe.py @@ -0,0 +1,42 @@ +"""Universe:选股/回测的股票范围执行器(ARCHITECTURE_v2 §14/§20 Universe 输入)。 + +把 ResearchService.filter_stocks 的语义规则化并集中于此: +- 当前日与历史日(as_of)都必须正确:退市股(delist < as_of)、上市时间(list_date) +- exclude_st 按**当前名称快照**含 ST 判定(历史可追溯数据;历史改名无法回溯,属近似, + 见结果 unimplemented 说明) +- exclude_suspended 依赖停牌数据表(尚未建模),此处不做剔除,由上层显式标注 +- symbols 白名单:非空时仅这些 symbol 参与(自选池 / 测试用) +""" + +from __future__ import annotations + +from collections.abc import Sequence +from datetime import date + +from app.domain.entities.market import Stock +from app.domain.entities.research import UniverseSpec + + +def filter_stocks( + stocks: Sequence[Stock], + universe: UniverseSpec, + as_of: date, +) -> list[Stock]: + """按股票池口径过滤,返回 as_of 时点应纳入的股票列表。""" + symbols = set(universe.symbols) if universe.symbols else None + out: list[Stock] = [] + for s in stocks: + if symbols is not None and s.symbol not in symbols: + continue + if s.delist_date is not None and s.delist_date < as_of: + continue + if universe.exclude_st and s.name and "ST" in s.name.upper(): + continue + if ( + universe.min_listing_days + and s.list_date + and (as_of - s.list_date).days < universe.min_listing_days + ): + continue + out.append(s) + return out diff --git a/backend/tests/test_universe.py b/backend/tests/test_universe.py new file mode 100644 index 0000000..7ee3217 --- /dev/null +++ b/backend/tests/test_universe.py @@ -0,0 +1,70 @@ +"""M6.1 UniverseFilter 测试:as_of 当前/历史日语义、ST、上市天数、退市、symbols 白名单。 + +filter_stocks 从 quant.universe 引入(原 quant.service 语义,规则化集中)。 +""" + +from __future__ import annotations + +from datetime import date + +from app.domain.entities.market import Stock +from app.domain.entities.research import UniverseSpec +from app.quant.universe import filter_stocks + + +def _stocks() -> list[Stock]: + return [ + Stock(symbol="600000.SH", name="正常股份", list_date=date(2000, 1, 1)), + Stock(symbol="600001.SH", name="ST 风险股份", list_date=date(2000, 1, 1)), + Stock(symbol="600002.SH", name="次新股", list_date=date(2024, 10, 1)), + Stock(symbol="600003.SH", name="已退市股", list_date=date(1995, 1, 1), + delist_date=date(2023, 6, 30)), + Stock(symbol="600004.SH", name="老股", list_date=date(1999, 1, 1)), + ] + + +def _sym(rows: list[Stock]) -> set[str]: + return {s.symbol for s in rows} + + +class TestUniverseFilter: + def test_current_day(self) -> None: + rows = filter_stocks(_stocks(), UniverseSpec(), as_of=date(2025, 1, 1)) + # ST、退市被剔除;次新股(上市<250 自然日)被剔除 + assert _sym(rows) == {"600000.SH", "600004.SH"} + + def test_historical_as_of_keeps_not_yet_delisted(self) -> None: + rows = filter_stocks(_stocks(), UniverseSpec(), as_of=date(2023, 1, 1)) + # 2023-01 时 600003 尚未退市(2023-06 退市)→ 应纳入 + assert "600003.SH" in _sym(rows) + assert "600002.SH" not in _sym(rows) # 2024-10 才上市,2023-01 尚不存在(上市天数不足被滤) + + def test_delisted_before_as_of_excluded(self) -> None: + rows = filter_stocks(_stocks(), UniverseSpec(exclude_st=False), as_of=date(2024, 1, 1)) + assert "600003.SH" not in _sym(rows) # 2023-06 已退市 + + def test_exclude_st_flag(self) -> None: + rows = filter_stocks(_stocks(), UniverseSpec(exclude_st=False), as_of=date(2025, 1, 1)) + assert "600001.SH" in _sym(rows) + rows2 = filter_stocks(_stocks(), UniverseSpec(exclude_st=True), as_of=date(2025, 1, 1)) + assert "600001.SH" not in _sym(rows2) + + def test_min_listing_days_zero_disables(self) -> None: + rows = filter_stocks( + _stocks(), UniverseSpec(min_listing_days=0, exclude_st=True), as_of=date(2025, 1, 1) + ) + assert "600002.SH" in _sym(rows) + + def test_symbols_whitelist(self) -> None: + rows = filter_stocks( + _stocks(), + UniverseSpec(symbols=["600000.SH", "600003.SH"]), + as_of=date(2025, 1, 1), + ) + # 白名单内的 ST/退市过滤仍然生效:600003 已退市被滤,仅剩 600000 + assert _sym(rows) == {"600000.SH"} + + def test_empty_whitelist_means_all(self) -> None: + assert UniverseSpec().symbols == [] + rows = filter_stocks(_stocks(), UniverseSpec(symbols=[]), as_of=date(2025, 1, 1)) + assert "600004.SH" in _sym(rows)