diff --git a/backend/app/application/services/selection_service.py b/backend/app/application/services/selection_service.py index dde9052..ee8d84b 100644 --- a/backend/app/application/services/selection_service.py +++ b/backend/app/application/services/selection_service.py @@ -2,10 +2,11 @@ - 输入:SelectionQuery(universe + method + factors/conditions + top_n/pct + as_of) - 装配:股票池(universe 过滤)→ 行情长表(含预热窗口)→ Selection Engine -- 输出:SelectionResult(可解释:factor_values / selection_reason) -- 未来函数红线:全部数据只取到 <= as_of(v2 §9);财务条件(后续)只取已公告值 + (method=score 因子评分 / method=condition 结构化条件) +- 输出:SelectionResult(可解释:factor_values / filter_status / selection_reason) +- 未来函数红线:行情只取 <= as_of;财务条件只取 announce_date <= as_of 的已公告值(v2 §9) -MVP 为同步执行(单日全市场因子计算量轻);如需异步可复用 Job 链路(M6.3 决策)。 +MVP 为同步执行(单日全市场因子/条件计算量轻);如需异步可复用 Job 链路。 """ from __future__ import annotations @@ -14,11 +15,23 @@ from datetime import date, timedelta import pandas as pd +from app.domain.entities.market import FinancialIndicator from app.domain.entities.selection import SelectionQuery, SelectionResult -from app.domain.repositories.market import DailyBarRepository, StockRepository -from app.quant.selection import factor_columns, run_score_selection +from app.domain.repositories.market import ( + DailyBarRepository, + FinancialRepository, + StockRepository, +) +from app.quant.selection import ( + condition_needed_columns, + factor_columns, + run_condition_selection, + run_score_selection, +) from app.quant.service import filter_stocks, load_daily_df +_FUNDAMENTAL_PREFIX = "fundamental." + class SelectionService: """选股用例入口:select(query) → SelectionResult(当前或历史 as_of)。""" @@ -27,17 +40,22 @@ class SelectionService: self, stock_repo: StockRepository, daily_repo: DailyBarRepository, + financial_repo: FinancialRepository | None = None, ) -> None: self._stock_repo = stock_repo self._daily_repo = daily_repo + self._financial_repo = financial_repo def select(self, query: SelectionQuery) -> SelectionResult: as_of = query.as_of or date.today() stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of) if not stocks: - return run_score_selection(pd.DataFrame(), query, as_of) + return self._run(query, pd.DataFrame(), stocks, as_of, financial={}) symbols = [s.symbol for s in stocks] - columns = sorted(factor_columns(query)) if query.method == "score" else ["close"] + if query.method == "score": + columns = sorted(factor_columns(query)) + else: + columns = sorted(condition_needed_columns(query)) daily = load_daily_df( self._daily_repo, symbols, @@ -45,6 +63,53 @@ class SelectionService: as_of, columns, ) - if daily.empty: + financial: dict[str, FinancialIndicator] = {} + if query.method == "condition" and self._uses_fundamental(query): + financial = self._load_financial(symbols, as_of) + return self._run(query, daily, stocks, as_of, financial) + + # ---- 内部 ---- + + def _run( + self, + query: SelectionQuery, + daily: pd.DataFrame, + stocks: list, + as_of: date, + financial: dict[str, FinancialIndicator], + ) -> SelectionResult: + if query.method == "score": return run_score_selection(daily, query, as_of) - return run_score_selection(daily, query, as_of) + return run_condition_selection(daily, stocks, query, as_of, financial) + + @staticmethod + def _uses_fundamental(query: SelectionQuery) -> bool: + for c in query.conditions: + if c.field.startswith(_FUNDAMENTAL_PREFIX) or ( + c.ref is not None and c.ref.startswith(_FUNDAMENTAL_PREFIX) + ): + return True + return False + + def _load_financial( + self, symbols: list[str], as_of: date + ) -> dict[str, FinancialIndicator]: + """按 announce_date <= as_of 批量取财务,每 symbol 保留最新一版。""" + if self._financial_repo is None: + raise ValueError("condition 引用了 fundamental.* 字段,但未注入 FinancialRepository") + getter = getattr(self._financial_repo, "list_announced_many", None) + if getter is not None: + rows = list(getter(symbols, as_of)) + else: # 回退逐只 + rows = [] + for sym in symbols: + rows.extend(self._financial_repo.list_announced(sym, as_of)) + by_symbol: dict[str, FinancialIndicator] = {} + for row in rows: + cur = by_symbol.get(row.symbol) + if cur is None or (row.announce_date, row.report_date) > ( + cur.announce_date, + cur.report_date, + ): + by_symbol[row.symbol] = row + return by_symbol diff --git a/backend/app/domain/entities/selection.py b/backend/app/domain/entities/selection.py index a6c4583..d5c575b 100644 --- a/backend/app/domain/entities/selection.py +++ b/backend/app/domain/entities/selection.py @@ -43,12 +43,13 @@ class SelectionQuery(BaseModel): @model_validator(mode="after") def _check_method_args(self) -> SelectionQuery: - if self.method == "score" and not self.factors: - raise ValueError("method=score 需要至少一个 factors") + if self.method == "score": + if not self.factors: + raise ValueError("method=score 需要至少一个 factors") + if self.top_n is None and self.top_pct is None: + raise ValueError("method=score 需要 top_n 与 top_pct 至少提供一个") if self.method == "condition" and not self.conditions: raise ValueError("method=condition 需要至少一个 conditions") - if self.top_n is None and self.top_pct is None: - raise ValueError("top_n 与 top_pct 至少提供一个") return self @model_validator(mode="after") diff --git a/backend/app/domain/repositories/market.py b/backend/app/domain/repositories/market.py index 35fda1f..f9c2e94 100644 --- a/backend/app/domain/repositories/market.py +++ b/backend/app/domain/repositories/market.py @@ -86,6 +86,17 @@ class FinancialRepository(Protocol): ) -> list[FinancialIndicator]: """只返回 announce_date <= as_of_date 的记录 —— 未来函数红线。""" + def list_announced_many( + self, + symbols: Sequence[str], + as_of_date: date, + ) -> list[FinancialIndicator]: + """批量版:返回这些股票 announce_date <= as_of_date 的全部记录。 + + 供选股/截面研究一次性取财务字段(调用方按需取每 symbol 最新一版)。 + 实现可选 —— 未提供时 SelectionService 回退逐只 list_announced。 + """ + class SyncLogRepository(Protocol): def add(self, log: SyncLog) -> SyncLog: ... diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py index 8b8608b..a417982 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/market_impl.py @@ -280,6 +280,28 @@ class SqlAlchemyFinancialRepository: rows = self._session.scalars(stmt).all() return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows] + def list_announced_many( + self, + symbols: Sequence[str], + as_of_date: date, + ) -> list[FinancialIndicator]: + """批量:这些股票 announce_date <= as_of_date 的全部记录(防未来函数)。""" + if not symbols: + return [] + rows = self._session.scalars( + select(FinancialIndicatorModel) + .where( + FinancialIndicatorModel.symbol.in_(list(symbols)), + FinancialIndicatorModel.announce_date <= as_of_date, + ) + .order_by( + FinancialIndicatorModel.symbol, + FinancialIndicatorModel.announce_date, + FinancialIndicatorModel.report_date, + ) + ).all() + return [FinancialIndicator.model_validate(r, from_attributes=True) for r in rows] + class SqlAlchemySyncLogRepository: def __init__(self, session: Session) -> None: diff --git a/backend/app/quant/selection.py b/backend/app/quant/selection.py index e232d9f..9f65c7e 100644 --- a/backend/app/quant/selection.py +++ b/backend/app/quant/selection.py @@ -14,6 +14,7 @@ from datetime import date import pandas as pd +from app.domain.entities.market import FinancialIndicator from app.domain.entities.selection import ( SelectionCandidate, SelectionQuery, @@ -163,6 +164,195 @@ def _symbol_count(daily: pd.DataFrame) -> int: return int(daily["symbol"].nunique()) if not daily.empty and "symbol" in daily else 0 +# ---------- method=condition:结构化条件选股(M6.2) ---------- + +# 技术字段:预计算派生量 + 行情原列(原列需在装配列中才可用) +_TECH_DERIVED = ("ma20", "ma60") +_STATIC_PREFIX = "static." +_FUNDAMENTAL_PREFIX = "fundamental." + + +def condition_needed_columns(query: SelectionQuery) -> set[str]: + """条件引用的行情列(fundamental/static 走元数据与财务表,不需要行情列)。""" + needed = {"close"} + names = [c.field for c in query.conditions] + [ + c.ref for c in query.conditions if c.ref and not c.ref.startswith(_FUNDAMENTAL_PREFIX) + ] + for f in names: + if not f or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)): + continue + if f in {"open", "high", "low", "close", "volume", "amount", *_TECH_DERIVED}: + if f not in _TECH_DERIVED: + needed.add(f) + continue + try: # 其余按已注册因子处理 + defn, _fn = get_factor(f) + except FactorError: + raise ValueError( + f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/static.*/fundamental.*)" + ) from None + needed.update(defn.requires) + return needed + + +def run_condition_selection( + daily: pd.DataFrame, + stocks: list, + query: SelectionQuery, + as_of: date | None, + financial: dict[str, FinancialIndicator] | None = None, +) -> SelectionResult: + """条件选股(v2 §14.1A):全部条件 AND 通过者入选(无排序;truncation 不适用)。 + + fields 域:static.*(股票基础)、close/volume/amount/ma20/ma60/已注册因子(行情)、 + fundamental.*(announce_date <= as_of 的最新已公告财务值 —— 防未来函数由 Service 取数保证)。 + """ + if query.method != "condition": + raise ValueError(f"run_condition_selection 需要 method=condition,当前 {query.method}") + obs = resolve_observation_date(daily, as_of) + resolved = (obs.date() if obs is not None else as_of) or date.today() + if obs is None or daily.empty: + return SelectionResult( + as_of_date=resolved, method=query.method, + statistics=SelectionStatistics(), candidates=[], + unimplemented=list(_UNIMPLEMENTED_DEFAULT), + config_snapshot=query.model_dump(mode="json"), + ) + + view = daily[pd.to_datetime(daily["trade_date"]) <= obs] + close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index() + close.index = pd.to_datetime(close.index) + + # 技术字段面板(obs 行) + tech: dict[str, pd.Series] = {} + for col in ("close", "open", "high", "low", "volume", "amount"): + if col in view.columns and col != "close": + panel = view.pivot(index="trade_date", columns="symbol", values=col).sort_index() + panel.index = pd.to_datetime(panel.index) + tech[col] = panel.loc[obs] + tech["close"] = close.loc[obs] + tech["ma20"] = close.rolling(20).mean().loc[obs] + tech["ma60"] = close.rolling(60).mean().loc[obs] + # 因子字段按需计算 + for cond in query.conditions: + for f in (cond.field, cond.ref): + if f is None or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)) or f in tech: + continue + if f in _TECH_DERIVED or f in ("close", "open", "high", "low", "volume", "amount"): + continue + try: + _defn, panel = compute_factor(f, view) + except FactorError: + continue # 已在 condition_needed_columns 报错;此处防御 + if obs in panel.index: + tech[f] = panel.loc[obs] + + statics = {s.symbol: s.model_dump() for s in stocks} + candidates: list[SelectionCandidate] = [] + passed_symbols: list[str] = [] + for sym in sorted(statics): + statuses: list[str] = [] + all_ok = True + for cond in query.conditions: + ok = _eval_condition(cond, sym, statics, tech, financial or {}) + statuses.append(f"{cond.field} {cond.op} {cond.ref or cond.value}: {'通过' if ok else '未通过'}") + all_ok = all_ok and ok + if all_ok: + passed_symbols.append(sym) + candidates.append( + SelectionCandidate( + symbol=sym, + rank=0, # 占位,末尾统一编号 + score=1.0, + filter_status=statuses, + selection_reason=[f"通过全部 {len(query.conditions)} 条条件"], + ) + ) + for rank, c in enumerate(candidates, start=1): + c.rank = rank + + return SelectionResult( + as_of_date=resolved, + method=query.method, + statistics=SelectionStatistics( + universe_size=len(statics), + evaluated=len(statics), + selected=len(candidates), + ), + candidates=candidates, + unimplemented=list(_UNIMPLEMENTED_DEFAULT) + [ + "条件选股为纯过滤(AND),未排序/未截断;如需排序请在 factors 中提供评分", + ], + config_snapshot=query.model_dump(mode="json"), + ) + + +def _eval_condition( + cond, + sym: str, + statics: dict, + tech: dict[str, pd.Series], + financial: dict, +) -> bool: + """求值单条条件:value 与 ref 二选一;left 与 right 同为 field 或 field vs 字面量。""" + left = _field_value(cond.field, sym, statics, tech, financial) + if cond.ref is not None: + right = _field_value(cond.ref, sym, statics, tech, financial) + else: + right = cond.value + return _compare(left, right, cond.op) + + +def _field_value(field, sym, statics, tech, financial): + if field.startswith(_STATIC_PREFIX): + return statics.get(sym, {}).get(field[len(_STATIC_PREFIX):]) + if field.startswith(_FUNDAMENTAL_PREFIX): + fin = financial.get(sym) + return getattr(fin, field[len(_FUNDAMENTAL_PREFIX):], None) if fin else None + series = tech.get(field) + if series is None: + return None + v = series.get(sym) + return None if v is None or (isinstance(v, float) and v != v) else v # NaN → None + + +def _compare(left, right, op: str) -> bool: + """混合比较:None 视为不可用 → 除 ne 外不通过;数值/字符串分别处理。""" + if op == "ne": + return left != right + if left is None or right is None: + return False + try: + if isinstance(left, (int, float)) or isinstance(right, (int, float)): + return _num_cmp(float(left), float(right), op) + except (TypeError, ValueError): + pass + # 字符串/其它:支持 eq/ne/in/not_in + if op == "eq": + return left == right + if op == "in": + return left in right + if op == "not_in": + return left not in right + if op in ("gt", "gte", "lt", "lte"): + return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False + return False + + +def _num_cmp(a: float, b: float, op: str) -> bool: + if op == "gt": + return a > b + if op == "gte": + return a >= b + if op == "lt": + return a < b + if op == "lte": + return a <= b + if op == "eq": + return a == b + return a != b + + def _to_float(v) -> float | None: if v is None: return None diff --git a/backend/tests/test_selection_condition.py b/backend/tests/test_selection_condition.py new file mode 100644 index 0000000..7a0ba14 --- /dev/null +++ b/backend/tests/test_selection_condition.py @@ -0,0 +1,228 @@ +"""M6.2 条件选股测试:结构化条件(static.*/tech 字段与因子/fundamental.*)与未来函数防护。 + +Fake(内存)Repository;财务行显式携带 announce_date 验证 as_of 可见性。 +""" + +from __future__ import annotations + +from datetime import date +from decimal import Decimal + +import pandas as pd +import pytest +from app.application.services.selection_service import SelectionService +from app.domain.entities.market import FinancialIndicator, Stock +from app.domain.entities.research import UniverseSpec +from app.domain.entities.selection import SelectionQuery + +from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily + +_SYMS = ["600000.SH", "600001.SH", "600002.SH"] + + +def _stocks(industries: dict[str, str] | None = None) -> list[Stock]: + industries = industries or {} + return [ + Stock( + symbol=s, name=f"股{i}", industry=industries.get(s), + list_date=date(1999, 1, 1), + ) + for i, s in enumerate(_SYMS) + ] + + +class _MemDailyRepo: + def __init__(self, df: pd.DataFrame) -> None: + self._bars_all = bars_dataframe_to_daily_bars(df) + + 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): + syms = set(symbols) + return [b for b in self._bars_all if b.symbol in syms and start <= b.trade_date <= end] + + def latest_date(self, symbol): + rows = [b.trade_date for b in self._bars_all if b.symbol == symbol] + return max(rows) if rows else None + + +class _MemFinancialRepo: + def __init__(self, rows: list[FinancialIndicator]) -> None: + self._rows = rows + + def list_announced(self, symbol, as_of_date, report_start=None): + return [ + r for r in self._rows + if r.symbol == symbol and r.announce_date <= as_of_date + and (report_start is None or r.report_date >= report_start) + ] + + def list_announced_many(self, symbols, as_of_date): + syms = set(symbols) + return [ + r for r in self._rows + if r.symbol in syms and r.announce_date <= as_of_date + ] + + +def _svc(df, stocks=None, fin_rows=None) -> SelectionService: + return SelectionService( + _MemStockRepo(stocks or _stocks()), + _MemDailyRepo(df), + _MemFinancialRepo(fin_rows or []), + ) + + +class _MemStockRepo: + def __init__(self, stocks) -> None: + self._stocks = stocks + + def list(self): + return self._stocks + + def get_by_symbol(self, symbol): + return next((s for s in self._stocks if s.symbol == symbol), None) + + +def _cond_query(conditions, **kw) -> SelectionQuery: + base = dict(method="condition", conditions=conditions, as_of=date(2024, 12, 31)) + base.update(kw) + return SelectionQuery(**base) + + +class TestStaticCondition: + def test_industry_filter(self) -> None: + df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320) + industries = {_SYMS[0]: "白酒", _SYMS[1]: "银行", _SYMS[2]: "白酒"} + svc = _svc(df, stocks=_stocks(industries)) + res = svc.select( + _cond_query( + [{"field": "static.industry", "op": "in", "value": ["白酒"]}], + universe=UniverseSpec(min_listing_days=0), + ) + ) + got = {c.symbol for c in res.candidates} + assert got == {_SYMS[0], _SYMS[2]} + assert res.statistics.selected == 2 + # 每候选带条件状态与理由 + assert all(c.filter_status for c in res.candidates) + assert all(c.selection_reason for c in res.candidates) + + def test_static_ne(self) -> None: + df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320) + industries = {_SYMS[0]: "白酒", _SYMS[1]: "银行", _SYMS[2]: "白酒"} + svc = _svc(df, stocks=_stocks(industries)) + res = svc.select( + _cond_query( + [{"field": "static.industry", "op": "ne", "value": "白酒"}], + universe=UniverseSpec(min_listing_days=0), + ) + ) + assert {c.symbol for c in res.candidates} == {_SYMS[1]} + + +class TestTechCondition: + def test_momentum_gt_zero(self) -> None: + df = synthetic_daily({_SYMS[0]: 0.008, _SYMS[1]: -0.008, _SYMS[2]: 0.002}, n=320) + svc = _svc(df) + res = svc.select( + _cond_query( + [{"field": "momentum_60", "op": "gt", "value": 0}], + universe=UniverseSpec(min_listing_days=0), + ) + ) + got = {c.symbol for c in res.candidates} + assert _SYMS[1] not in got # 下跌股 60 日动量为负 + assert _SYMS[0] in got and _SYMS[2] in got + + def test_close_above_ma60_ref(self) -> None: + df = synthetic_daily({_SYMS[0]: 0.006, _SYMS[1]: -0.006, _SYMS[2]: 0.0005}, n=320) + svc = _svc(df) + res = svc.select( + _cond_query( + [{"field": "close", "op": "gt", "ref": "ma60"}], + universe=UniverseSpec(min_listing_days=0), + ) + ) + got = {c.symbol for c in res.candidates} + assert _SYMS[1] not in got # 下跌股收盘在 MA60 之下 + assert _SYMS[0] in got + + def test_lte_threshold(self) -> None: + df = synthetic_daily({_SYMS[0]: 0.01, _SYMS[1]: -0.01, _SYMS[2]: 0.0001}, n=320) + svc = _svc(df) + res = svc.select( + _cond_query( + [{"field": "momentum_60", "op": "lte", "value": 0}], + universe=UniverseSpec(min_listing_days=0), + ) + ) + got = {c.symbol for c in res.candidates} + assert _SYMS[1] in got + assert _SYMS[0] not in got + + +class TestFundamentalCondition: + def _fin_rows(self) -> list[FinancialIndicator]: + return [ + FinancialIndicator( + symbol=_SYMS[0], report_date=date(2024, 9, 30), announce_date=date(2024, 10, 25), + eps=Decimal("3.5"), roe=Decimal("20.0"), source="tushare", + ), + # B:只在 as_of 之后才公告(未来数据)→ as_of 时不可见 + FinancialIndicator( + symbol=_SYMS[1], report_date=date(2024, 9, 30), announce_date=date(2025, 3, 30), + roe=Decimal("99.0"), source="tushare", + ), + # C:roe 低于阈值 + FinancialIndicator( + symbol=_SYMS[2], report_date=date(2024, 6, 30), announce_date=date(2024, 8, 20), + roe=Decimal("5.0"), source="tushare", + ), + ] + + def test_roe_filter_no_future_leak(self) -> None: + df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320) + svc = _svc(df, fin_rows=self._fin_rows()) + res = svc.select( + _cond_query( + [{"field": "fundamental.roe", "op": "gte", "value": 15}], + universe=UniverseSpec(min_listing_days=0), + ), + # 上面 helper 已带 as_of + ) + got = {c.symbol for c in res.candidates} + # A 可见且 roe=20 → 入选;B 公告在未来(防未来函数)→ 不入选;C roe=5 → 不入选 + assert got == {_SYMS[0]} + + def test_missing_financial_repo_raises(self) -> None: + df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320) + svc = SelectionService(_MemStockRepo(_stocks()), _MemDailyRepo(df)) # 无财务 repo + with pytest.raises(ValueError): + svc.select( + _cond_query( + [{"field": "fundamental.roe", "op": "gte", "value": 15}], + universe=UniverseSpec(min_listing_days=0), + ) + ) + + def test_as_of_earlier_excludes_announced_after(self) -> None: + """更早 as_of:A 的 roe=20 若在 as_of 之后才公告也不可见。""" + df = synthetic_daily({s: 0.002 for s in _SYMS}, n=320) + rows = [ + FinancialIndicator( + symbol=_SYMS[0], report_date=date(2024, 6, 30), + announce_date=date(2024, 10, 1), roe=Decimal("99.0"), source="tushare", + ) + ] + svc = _svc(df, fin_rows=rows) + res = svc.select( + SelectionQuery( + method="condition", + conditions=[{"field": "fundamental.roe", "op": "gte", "value": 50}], + as_of=date(2024, 9, 1), # announce(10-01) 尚未来 + universe=UniverseSpec(min_listing_days=0), + ) + ) + assert {c.symbol for c in res.candidates} == set() # 无人可见 roe → 全部不通过