- 首页双入口(智能选股/策略回测):引入 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>
371 lines
15 KiB
Python
371 lines
15 KiB
Python
"""选股执行引擎: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,
|
||
}
|