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>
This commit is contained in:
370
backend/app/screener/engine.py
Normal file
370
backend/app/screener/engine.py
Normal file
@@ -0,0 +1,370 @@
|
||||
"""选股执行引擎: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,
|
||||
}
|
||||
Reference in New Issue
Block a user