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:
2026-08-14 14:49:53 +08:00
parent e0b5228008
commit 528357c3f5
29 changed files with 1765 additions and 12 deletions

View 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,
}