Files
stock/backend/app/screener/engine.py
cirry 528357c3f5 feat: AI 自然语言选股(GLM)+ 全市场数据管道 + 远程 PostgreSQL
- 首页双入口(智能选股/策略回测):引入 vue-router,顶部导航
- 智能选股:自然语言 -> LLM 解析结构化条件(智谱 GLM,OpenAI 兼容,/v4 兼容)-> SQL 快照预筛 + pandas 指标过滤(复用 indicators 单一事实源)
- 条件模型:指标 vs 常数/指标(value_indicator,如 DIF>DEA、close<布林下轨)、lookback+match 表达连续N天/近N天任一天、市值/PE/PB/换手率快照条件、默认排除 ST/退市/北交所
- 全市场数据同步:按 trade_date 批量拉取未复权日线(与回测 candles qfq 隔离),交易日历/股票列表本地缓存,daily_basic 仅最新截面,Tushare 限频兜底(分钟级重试/小时级降级)
- 存储:DATABASE_URL 切远程 PostgreSQL(cirry.cn/stock),本地 SQLite 已移除
- .env 入库(私有仓库);smoke_test 扩展选股链路

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-14 14:49:53 +08:00

371 lines
15 KiB
Python
Raw Permalink 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 快照预筛缩小范围 -> 逐股指标计算过滤。
性能:预筛在 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,
}