Files
stock/backend/app/screener/engine.py
2026-08-16 20:20:59 +08:00

418 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""选股执行引擎SQL 快照预筛缩小范围 -> 逐股指标计算过滤。
性能:预筛在数据库索引上完成(毫秒级);指标阶段候选集通常数百~数千只 × ~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 ..data.symbols import plain_code
from ..models import Candle, DailySnapshot, 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), # 实际取 candles 最新 bar无需快照表
}
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 = Candle.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:
"""最新交易日截面预筛candles(全量不复权底座) + stock_basic + daily_snapshot。
只取 target_date 当日有交易的股票(停牌股无当日 bar自然排除与原先 market_daily
的 trade_date == target_date 行为一致。pct_chg 用前一交易日收盘价计算。
"""
snap_date = await session.scalar(select(func.max(DailySnapshot.trade_date))) or target_date
stmt = (
select(
Candle.symbol, StockBasic.ts_code, StockBasic.name, Candle.close,
DailySnapshot.total_mv, DailySnapshot.circ_mv,
DailySnapshot.pe_ttm, DailySnapshot.pb, DailySnapshot.turnover_rate,
)
.join(StockBasic, StockBasic.symbol == Candle.symbol)
.outerjoin(DailySnapshot, and_(DailySnapshot.ts_code == StockBasic.ts_code,
DailySnapshot.trade_date == snap_date))
.where(Candle.timeframe == "1d", Candle.ts == 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_(StockBasic.ts_code.like("%.BJ")))
for c in conds.snapshot:
stmt = stmt.where(_snapshot_clause(c))
rows = (await session.execute(stmt)).all()
df = pd.DataFrame(rows, columns=["symbol", "ts_code", "name", "close",
"total_mv", "circ_mv", "pe_ttm", "pb", "turnover_rate"])
if df.empty:
return df[_CAND_COLS]
# pct_chgcandles 无现成涨跌幅列,用前一交易日的收盘价计算
prev_dt = await session.scalar(
select(func.max(Candle.ts)).where(Candle.timeframe == "1d", Candle.ts < target_date)
)
prev_map: dict[str, float] = {}
if prev_dt is not None:
pr = await session.execute(
select(Candle.symbol, Candle.close).where(
Candle.timeframe == "1d", Candle.ts == prev_dt,
Candle.symbol.in_(df["symbol"].tolist()),
)
)
prev_map = {r.symbol: r.close for r in pr}
prev = df["symbol"].map(prev_map)
def _pct(c, p) -> float | None:
if p is None or p != p or float(p) == 0:
return None
return (float(c) / float(p) - 1) * 100
df["pct_chg"] = [_pct(c, p) for c, p in zip(df["close"], prev)]
return df[_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 线窗口candles 全量不复权底座)。
候选 <= 2000 用 IN 精确圈定;否则拉全窗口再 pandas 过滤。
ts_code 按候选集映射回 symbol 查询pct_chg 用每股收盘价环比计算。
"""
symbols = [plain_code(t) for t in ts_codes]
stmt = select(
Candle.symbol, Candle.ts, Candle.open, Candle.high, Candle.low, Candle.close,
).where(Candle.timeframe == "1d", Candle.ts >= min_date, Candle.ts <= target_date)
if len(symbols) <= 2000:
stmt = stmt.where(Candle.symbol.in_(set(symbols)))
rows = (await session.execute(stmt)).all()
df = pd.DataFrame(rows, columns=["symbol", "trade_date", "open", "high", "low", "close"])
if not df.empty and len(symbols) > 2000:
df = df[df["symbol"].isin(set(symbols))]
df = df.sort_values(["symbol", "trade_date"]).reset_index(drop=True)
if df.empty:
df["ts_code"] = pd.Series(dtype=object)
df["pct_chg"] = pd.Series(dtype=object)
else:
df["pct_chg"] = df.groupby("symbol")["close"].pct_change() * 100
ts_map = {plain_code(t): t for t in ts_codes}
df["ts_code"] = df["symbol"].map(ts_map)
return df[["ts_code", "trade_date", "open", "high", "low", "close", "pct_chg"]]
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(Candle.ts)).where(Candle.timeframe == "1d")
)
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 线窗口:按 candles 全量交易日序列回溯 needed 根(不再受同步窗口限制)
need = _max_needed_bars(conds)
dates_res = await session.execute(
select(Candle.ts).where(Candle.timeframe == "1d").distinct()
.order_by(Candle.ts.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,
}