"""选股执行引擎:SQL 快照预筛缩小范围 -> 逐股指标计算过滤。 性能:预筛在 SQLite 索引上完成(毫秒级);指标阶段候选集通常数百~数千只 × ~90 根 bar, pandas 逐股计算(复用 app/indicators,指标按 (族, 参数) 去重计算),秒级完成。 """ from __future__ import annotations from dataclasses import dataclass from datetime import datetime from typing import Callable import pandas as pd from sqlalchemy import and_, func, not_, or_, select from sqlalchemy.ext.asyncio import AsyncSession from .. import indicators as ind from ..models import DailySnapshot, MarketDaily, StockBasic from ..schemas import IndicatorCondition, ScreenConditions, SnapshotCondition # 快照字段 -> (DB 列, 中文标签, LLM 值 -> DB 值的换算乘数) # 单位约定:LLM 市值亿元 -> DB 万元 ×1e4;其余(pe/pb/换手率%)直存 SNAPSHOT_FIELDS: dict[str, tuple[str, str, float]] = { "total_mv": ("total_mv", "总市值(亿)", 1e4), "circ_mv": ("circ_mv", "流通市值(亿)", 1e4), "pe_ttm": ("pe_ttm", "市盈率TTM", 1.0), "pb": ("pb", "市净率", 1.0), "turnover_rate": ("turnover_rate", "换手率%", 1.0), "close": ("close", "最新价", 1.0), # 实际取 market_daily.close,无需快照表 } class DataNotReadyError(RuntimeError): """全市场数据未同步/不完整,选股无法执行。""" # ---------- 指标族注册表(复用 app/indicators,单一事实源) ---------- def _kdj(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: out = ind.kdj(df["high"], df["low"], df["close"], n=int(p["n"]), m1=int(p["m1"]), m2=int(p["m2"])) return {"kdj_k": out["k"], "kdj_d": out["d"], "kdj_j": out["j"]} def _macd(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: out = ind.macd(df["close"], fast=int(p["fast"]), slow=int(p["slow"]), signal=int(p["signal"])) return {"macd_dif": out["macd"], "macd_dea": out["signal"], "macd_hist": out["hist"]} def _rsi(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: return {"rsi": ind.rsi(df["close"], period=int(p["period"]))} def _ma(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: return {"ma": ind.ma(df["close"], period=int(p["period"]))} def _boll(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: out = ind.bollinger(df["close"], period=int(p["period"]), std=float(p["std"])) return {"boll_upper": out["upper"], "boll_mid": out["mid"], "boll_lower": out["lower"]} def _close(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: return {"close": df["close"]} def _pct_chg(df: pd.DataFrame, p: dict) -> dict[str, pd.Series]: return {"pct_chg": df["pct_chg"]} @dataclass(frozen=True) class FamilyDef: label: str # 族中文标签(条件回显/表头) param_names: tuple[str, ...] # 接受的参数名 defaults: dict[str, float] min_bars: int # 指标可信所需最小 bar 数(近似) compute: Callable[[pd.DataFrame, dict], dict[str, pd.Series]] FAMILIES: dict[str, FamilyDef] = { "kdj": FamilyDef("KDJ", ("n", "m1", "m2"), {"n": 9, "m1": 3, "m2": 3}, 30, _kdj), "macd": FamilyDef("MACD", ("fast", "slow", "signal"), {"fast": 12, "slow": 26, "signal": 9}, 60, _macd), "rsi": FamilyDef("RSI", ("period",), {"period": 14}, 25, _rsi), "ma": FamilyDef("MA", ("period",), {"period": 20}, 25, _ma), "boll": FamilyDef("BOLL", ("period", "std"), {"period": 20, "std": 2}, 25, _boll), "close": FamilyDef("收盘价", (), {}, 1, _close), "pct_chg": FamilyDef("日涨跌幅", (), {}, 1, _pct_chg), } INDICATOR_FAMILY: dict[str, str] = { "kdj_k": "kdj", "kdj_d": "kdj", "kdj_j": "kdj", "macd_dif": "macd", "macd_dea": "macd", "macd_hist": "macd", "rsi": "rsi", "ma": "ma", "boll_upper": "boll", "boll_mid": "boll", "boll_lower": "boll", "close": "close", "pct_chg": "pct_chg", } _IND_SUFFIX = {"kdj_k": "K", "kdj_d": "D", "kdj_j": "J", "macd_dif": "DIF", "macd_dea": "DEA", "macd_hist": "柱", "boll_upper": "上轨", "boll_mid": "中轨", "boll_lower": "下轨"} def _family_of(indicator: str) -> str: fam = INDICATOR_FAMILY.get(indicator) if not fam: raise ValueError(f"未知指标: {indicator}(白名单见 screener/llm.py)") return fam def _params_for(indicator: str, params: dict, extra: dict | None = None) -> dict: """按指标族过滤参数:extra(如 value_params)优先,其次 params,最后默认值。""" fam = FAMILIES[_family_of(indicator)] merged = {**{k: v for k, v in params.items() if k in fam.param_names}, **{k: v for k, v in (extra or {}).items() if k in fam.param_names}} return {**fam.defaults, **merged} def _resolve_params(cond: IndicatorCondition, target: str) -> dict: """条件里 value_indicator 目标的参数:value_params > params(按目标族过滤)> 默认。""" return _params_for(target, cond.params, cond.value_params) def _params_str(p: dict) -> str: vals = [str(int(v)) if float(v).is_integer() else str(v) for v in p.values()] return f"({','.join(vals)})" if vals else "" def indicator_label(indicator: str, params: dict) -> str: """指标展示名,如 'KDJ J(9,3,3)' / 'MA(20)' / '收盘价'。""" fam = FAMILIES[_family_of(indicator)] suffix = _IND_SUFFIX.get(indicator) name = f"{fam.label} {suffix}".strip() if suffix else fam.label return name + _params_str(params) # ---------- 指标序列获取(按 (族, 参数) 去重计算) ---------- def _series_for(df_stock: pd.DataFrame, indicator: str, params: dict, cache: dict) -> pd.Series | None: """获取单股某指标序列。cache: {"_families": set[(族, 参数键)], (指标名, 参数键): Series}。 同族同参数只计算一次(如 KDJ 三个值共享一次计算);MA5/MA20 则按参数各自计算。 """ fam = _family_of(indicator) pk = tuple(sorted(params.items())) fk = (fam, pk) if fk not in cache["_families"]: for name, s in FAMILIES[fam].compute(df_stock, params).items(): cache[(name, pk)] = s cache["_families"].add(fk) return cache.get((indicator, pk)) def _op_mask(s: pd.Series, other: pd.Series, cond: IndicatorCondition) -> pd.Series: """按 op 生成布尔掩码(NaN -> False)。""" if cond.op == "gt": mask = s > other elif cond.op == "ge": mask = s >= other elif cond.op == "lt": mask = s < other elif cond.op == "le": mask = s <= other elif cond.op == "between": hi = cond.value2 if cond.value2 is not None else cond.value mask = (s >= cond.value) & (s <= hi) else: raise ValueError(f"未知 op: {cond.op}") return mask.fillna(False).astype(bool) def _eval_condition(df_stock: pd.DataFrame, cond: IndicatorCondition, cache: dict) -> tuple[bool, float | None]: """在单股 df 上判定条件;返回 (是否命中, 指标最新值)。数据不足视为不满足。""" fam = _family_of(cond.indicator) if len(df_stock) < max(FAMILIES[fam].min_bars, cond.lookback): return False, None s = _series_for(df_stock, cond.indicator, _params_for(cond.indicator, cond.params), cache) if s is None: return False, None if cond.value_indicator: target = _series_for(df_stock, cond.value_indicator, _resolve_params(cond, cond.value_indicator), cache) if target is None: return False, None else: target = pd.Series(cond.value, index=s.index) window = _op_mask(s, target, cond).tail(cond.lookback) hit = bool(window.all()) if cond.match == "all" else bool(window.any()) latest = s.iloc[-1] return hit, (None if latest != latest else float(latest)) # ---------- SQL 预筛 ---------- def _snapshot_clause(cond: SnapshotCondition): """单条快照条件 -> SQLAlchemy 表达式(快照缺失的股 NULL 比较自然不满足)。""" if cond.field not in SNAPSHOT_FIELDS: raise ValueError(f"未知快照字段: {cond.field}") col_name, _, scale = SNAPSHOT_FIELDS[cond.field] col = MarketDaily.close if cond.field == "close" else getattr(DailySnapshot, col_name) lo = cond.value * scale hi = (cond.value2 * scale) if cond.op == "between" and cond.value2 is not None else None if cond.op == "gt": return col > lo if cond.op == "ge": return col >= lo if cond.op == "lt": return col < lo if cond.op == "le": return col <= lo if hi is not None: return and_(col >= lo, col <= hi) raise ValueError(f"未知 op: {cond.op}") _CAND_COLS = ["ts_code", "name", "close", "pct_chg", "total_mv", "circ_mv", "pe_ttm", "pb", "turnover_rate"] async def _prefilter(session: AsyncSession, conds: ScreenConditions, target_date: datetime) -> pd.DataFrame: """最新交易日截面预筛:快照条件 + 排除项 + 名称/收盘价。返回候选 DataFrame。""" stmt = ( select( MarketDaily.ts_code, StockBasic.name, MarketDaily.close, MarketDaily.pct_chg, DailySnapshot.total_mv, DailySnapshot.circ_mv, DailySnapshot.pe_ttm, DailySnapshot.pb, DailySnapshot.turnover_rate, ) .join(StockBasic, StockBasic.ts_code == MarketDaily.ts_code) .outerjoin(DailySnapshot, and_(DailySnapshot.ts_code == MarketDaily.ts_code, DailySnapshot.trade_date == target_date)) .where(MarketDaily.trade_date == target_date) ) if conds.exclude_delisted: stmt = stmt.where(StockBasic.list_status == "L") if conds.exclude_st: stmt = stmt.where(not_(or_(StockBasic.name.like("%ST%"), StockBasic.name.like("%退%")))) if conds.exclude_bj: stmt = stmt.where(not_(MarketDaily.ts_code.like("%.BJ"))) for c in conds.snapshot: stmt = stmt.where(_snapshot_clause(c)) rows = (await session.execute(stmt)).all() return pd.DataFrame(rows, columns=_CAND_COLS) # ---------- 主流程 ---------- def _max_needed_bars(conds: ScreenConditions) -> int: """指标阶段需要的最大 bar 数(min_bars + lookback),用于圈定 K 线窗口。""" need = 1 for c in conds.indicator: need = max(need, FAMILIES[_family_of(c.indicator)].min_bars + c.lookback) if c.value_indicator: need = max(need, FAMILIES[_family_of(c.value_indicator)].min_bars + c.lookback) return need async def _load_bars(session: AsyncSession, ts_codes: list[str], target_date: datetime, min_date: datetime) -> pd.DataFrame: """载入候选股的 K 线窗口。候选 <= 2000 用 IN 精确圈定;否则拉全窗口再 pandas 过滤。""" stmt = select( MarketDaily.ts_code, MarketDaily.trade_date, MarketDaily.open, MarketDaily.high, MarketDaily.low, MarketDaily.close, MarketDaily.pct_chg, ).where(MarketDaily.trade_date >= min_date, MarketDaily.trade_date <= target_date) if len(ts_codes) <= 2000: stmt = stmt.where(MarketDaily.ts_code.in_(set(ts_codes))) rows = (await session.execute(stmt)).all() df = pd.DataFrame(rows, columns=["ts_code", "trade_date", "open", "high", "low", "close", "pct_chg"]) if not df.empty and len(ts_codes) > 2000: df = df[df["ts_code"].isin(set(ts_codes))] return df.sort_values(["ts_code", "trade_date"]).reset_index(drop=True) def _f(v) -> float | None: """None/NaN 安全转 float|None。""" if v is None: return None v = float(v) return None if v != v else v def _yi(v) -> float | None: """万元 -> 亿元(None/NaN 安全,保留两位)。""" v = _f(v) return None if v is None else round(v / 1e4, 2) def _item_from_row(row) -> dict: return { "ts_code": row["ts_code"], "name": row["name"], "close": _f(row["close"]), "pct_chg": _f(row["pct_chg"]), "total_mv": _yi(row["total_mv"]), "circ_mv": _yi(row["circ_mv"]), "pe_ttm": _f(row["pe_ttm"]), "pb": _f(row["pb"]), "turnover_rate": _f(row["turnover_rate"]), } async def run_screen(session: AsyncSession, conds: ScreenConditions, limit: int) -> dict: """主流程:预筛 -> 逐股指标过滤 -> 组装 items + indicator_labels。""" if not conds.indicator and not conds.snapshot: raise ValueError("筛选条件为空") target_date = await session.scalar(select(func.max(MarketDaily.trade_date))) if target_date is None: raise DataNotReadyError("全市场数据未同步:请先在选股页点击「同步市场数据」") if any(c.field != "close" for c in conds.snapshot): if await session.scalar(select(func.max(DailySnapshot.trade_date))) is None: raise DataNotReadyError( "每日指标数据缺失(可能是 Tushare 积分不足,daily_basic 接口不可用):" "市值/市盈率等条件无法使用,纯指标条件不受影响" ) cand = await _prefilter(session, conds, target_date) labels: dict[str, str] = {} for c in conds.indicator: labels.setdefault(c.indicator, indicator_label(c.indicator, _params_for(c.indicator, c.params))) if c.value_indicator: labels.setdefault(c.value_indicator, indicator_label(c.value_indicator, _resolve_params(c, c.value_indicator))) items: list[dict] = [] if cand.empty: pass elif not conds.indicator: # 纯快照条件:预筛结果即命中 items = [_item_from_row(row) | {"indicators": {}} for _, row in cand.iterrows()] else: # 圈定 K 线窗口:按已同步交易日序列回溯 needed 根 need = _max_needed_bars(conds) dates_res = await session.execute( select(MarketDaily.trade_date).distinct().order_by(MarketDaily.trade_date.desc()).limit(need) ) min_date = min(r[0] for r in dates_res) bars = await _load_bars(session, cand["ts_code"].tolist(), target_date, min_date) cand_rows = {r["ts_code"]: r for _, r in cand.iterrows()} for ts_code, g in bars.groupby("ts_code", sort=False): row = cand_rows.get(ts_code) if row is None: continue cache: dict = {"_families": set()} ok = True ind_values: dict[str, float | None] = {} for c in conds.indicator: hit, latest = _eval_condition(g, c, cache) if not hit: ok = False break if c.indicator not in ind_values: ind_values[c.indicator] = latest if c.value_indicator and c.value_indicator not in ind_values: s = cache.get((c.value_indicator, tuple(sorted(_resolve_params(c, c.value_indicator).items())))) ind_values[c.value_indicator] = None if s is None or s.iloc[-1] != s.iloc[-1] else float(s.iloc[-1]) if ok: items.append(_item_from_row(row) | {"indicators": ind_values}) # 默认总市值降序(缺失排最后),截断 limit items.sort(key=lambda x: (x["total_mv"] is None, -(x["total_mv"] or 0))) return { "conditions": conds, "trade_date": target_date, "total": len(items), "items": items[:limit], "indicator_labels": labels, }