diff --git a/.gitignore b/.gitignore index ec67911..bd87e5f 100644 --- a/.gitignore +++ b/.gitignore @@ -25,9 +25,7 @@ frontend/dist/ .cache/ # ---------- Env / Secrets ---------- -.env -.env.* -!.env.example +# .env 直接提交(私有仓库,配置含连接串即开即用) *.pem *.key diff --git a/backend/.env b/backend/.env new file mode 100644 index 0000000..4449133 --- /dev/null +++ b/backend/.env @@ -0,0 +1,10 @@ +DATABASE_URL=postgresql+asyncpg://postgres:Cirry0115@cirry.cn:5432/stock +TUSHARE_TOKEN=d0bc5620d6523ae40f379ed4415576f58dca2361f2f47a68cdcd0a98 +DATA_ADJUST=qfq +DATA_DEFAULT_START=20200101 + +# ---- LLM(智能选股;智谱 GLM,OpenAI 兼容协议)---- +# key 在 https://bigmodel.cn 控制台获取,格式形如 xxxxxxxx.yyyyyyyy(id.secret) +LLM_BASE_URL=https://open.bigmodel.cn/api/paas/v4 +LLM_API_KEY=ea24bbdd3d2d4dd2b8f03de4c9a5d984.9X1Hz1yKx0VKSnrU +LLM_MODEL=glm-5.2 diff --git a/backend/.env.example b/backend/.env.example index 0a365c8..a12d6b4 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -10,6 +10,22 @@ TUSHARE_TOKEN=你的token DATA_ADJUST=qfq # 复权:qfq 前复权 / hfq 后复权 / 留空不复权 DATA_DEFAULT_START=20200101 +# ---- LLM(智能选股;OpenAI 兼容协议,任选一家)---- +# 留空则智能选股不可用,其余功能不受影响。 + +# 智谱 GLM(key 在 https://bigmodel.cn 获取,格式 id.secret) +LLM_BASE_URL=https://open.bigmodel.cn/api/paas/v4 +LLM_API_KEY=你的key +LLM_MODEL=glm-5.2 + +# 或 DeepSeek:把上面三行换成 +# LLM_BASE_URL=https://api.deepseek.com +# LLM_MODEL=deepseek-chat + +# ---- 智能选股(可选覆盖)---- +# SCREENER_MARKET_DAYS=90 # 全市场同步窗口(交易日数) +# SCREENER_SYNC_INTERVAL=0.35 # 同步调用间隔(秒),Tushare 控频 + # ---- A股交易成本(基准日 2026-08,可覆盖;详见 app/commission.py)---- # STAMP_DUTY_RATE=0.0005 # 印花税 0.05%,单边卖出 # TRANSFER_FEE_RATE=0.00001 # 过户费 0.001%,沪深双边 diff --git a/backend/app/api.py b/backend/app/api.py index cd1b44a..4a04757 100644 --- a/backend/app/api.py +++ b/backend/app/api.py @@ -1,8 +1,11 @@ """HTTP 路由(OpenAPI 契约的载体)。 - GET /api/health 健康检查 - GET /api/candles/{sym} 取 K 线(支持 1d/1w/1M/1y 周期,日线为基底聚合) - POST /api/backtest 跑回测,返回 K线+指标+买卖点+净值+绩效 + GET /api/health 健康检查 + GET /api/candles/{sym} 取 K 线(支持 1d/1w/1M/1y 周期,日线为基底聚合) + POST /api/backtest 跑回测,返回 K线+指标+买卖点+净值+绩效 + POST /api/screener/run 智能选股:自然语言 -> 条件 -> 全市场筛选 + POST /api/screener/sync 启动全市场数据同步(后台任务) + GET /api/screener/sync/status 同步任务状态与数据实况 """ from __future__ import annotations @@ -14,6 +17,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from .backtest.engine import BacktestConfig, run_backtest from .backtest.strategies import build_strategy +from .config import settings from .data import fetcher, repository from .data.aggregation import bars_per_year, resample_bars from .data.synthetic import seed_if_empty @@ -27,10 +31,17 @@ from .schemas import ( EquityPoint, IndicatorOut, MetricsOut, + ScreenerRunRequest, + ScreenerRunResponse, + ScreenerSyncRequest, + ScreenerSyncStatus, SignalOut, SyncRequest, SyncResponse, ) +from .screener import engine, market_sync +from .screener.engine import DataNotReadyError +from .screener.llm import ScreenerError, parse_conditions router = APIRouter(prefix="/api") @@ -165,3 +176,46 @@ async def backtest( final_position=result["final_position"], initial_cash=req.initial_cash, ) + + +# ---------- 智能选股 ---------- +@router.post("/screener/run", response_model=ScreenerRunResponse) +async def screener_run( + req: ScreenerRunRequest, session: AsyncSession = Depends(get_session) +) -> ScreenerRunResponse: + """自然语言 -> LLM 解析条件 -> 全市场筛选。也可直传 conditions 跳过 LLM(微调再跑)。""" + try: + conds = req.conditions or await parse_conditions(req.text) + if not conds.indicator and not conds.snapshot: + raise HTTPException(status_code=400, detail="AI 未从描述中解析出任何筛选条件,请换种说法") + result = await engine.run_screen(session, conds, settings.screener_default_limit) + return ScreenerRunResponse(**result) + except HTTPException: + raise + except DataNotReadyError as e: + raise HTTPException(status_code=409, detail=str(e)) + except ValueError as e: # 未知指标/字段、条件为空 + raise HTTPException(status_code=400, detail=str(e)) + except ScreenerError as e: + code = 503 if "未配置 LLM_API_KEY" in str(e) else 502 + raise HTTPException(status_code=code, detail=str(e)) + + +@router.post("/screener/sync", response_model=ScreenerSyncStatus) +async def screener_sync_start( + req: ScreenerSyncRequest, session: AsyncSession = Depends(get_session) +) -> ScreenerSyncStatus: + """启动全市场数据同步(后台任务,立即返回状态)。""" + try: + await market_sync.start_sync(session, req.days, req.force) + except ScreenerError as e: + raise HTTPException(status_code=503, detail=str(e)) + status = await market_sync.get_sync_status(session) + return ScreenerSyncStatus(**{k: status.get(k) for k in ScreenerSyncStatus.model_fields}) + + +@router.get("/screener/sync/status", response_model=ScreenerSyncStatus) +async def screener_sync_status(session: AsyncSession = Depends(get_session)) -> ScreenerSyncStatus: + """同步任务状态 + 数据实况(最新交易日/行数/ready)。""" + status = await market_sync.get_sync_status(session) + return ScreenerSyncStatus(**{k: status.get(k) for k in ScreenerSyncStatus.model_fields}) diff --git a/backend/app/config.py b/backend/app/config.py index f16fc3c..7e65ec5 100644 --- a/backend/app/config.py +++ b/backend/app/config.py @@ -15,6 +15,17 @@ class Settings(BaseSettings): data_adjust: str = "qfq" # 复权:qfq 前复权 / hfq 后复权 / "" 不复权 data_default_start: str = "20200101" # 默认拉取起点(约近 5 年) + # ---- LLM(智能选股的自然语言解析;DeepSeek,OpenAI 兼容协议,可换任意兼容网关)---- + llm_base_url: str = "https://api.deepseek.com" + llm_api_key: str = "" # 留空则智能选股不可用(其余功能不受影响) + llm_model: str = "deepseek-chat" + llm_timeout: float = 60.0 + + # ---- 智能选股 ---- + screener_market_days: int = 90 # 全市场同步窗口(交易日数) + screener_default_limit: int = 200 # 选股结果条数上限 + screener_sync_interval: float = 0.35 # 全市场批量调用间隔(秒),Tushare 控频 + # A股交易成本(基准日 2026-08)——做成可配置参数,便于将来按生效日期版本化 stamp_duty_rate: float = 0.0005 # 印花税 0.05%,单边卖出(2023-08-28 减半) transfer_fee_rate: float = 0.00001 # 过户费 0.001%,沪深双边(2022 调整) diff --git a/backend/app/models.py b/backend/app/models.py index 73fe066..6146b08 100644 --- a/backend/app/models.py +++ b/backend/app/models.py @@ -3,6 +3,9 @@ Candle 表设计与 TimescaleDB hypertable 完全兼容:将来在目标 PG 库执行 SELECT create_hypertable('candles', 'ts'); 即可升级为时序表 + Continuous Aggregates 多周期预聚合,无需改表结构。 + +智能选股三表(stock_basic / market_daily / daily_snapshot)与回测 candles(qfq) +完全隔离:选股用未复权日线按 trade_date 全市场批量落地,避免污染回测复权缓存。 """ from datetime import datetime @@ -55,3 +58,75 @@ class BacktestRun(Base): max_drawdown: Mapped[float] = mapped_column(Float, default=0.0) sharpe: Mapped[float] = mapped_column(Float, default=0.0) num_trades: Mapped[int] = mapped_column(Integer, default=0) + + +class StockBasic(Base): + """A股股票列表(stock_basic 快照;选股展示名称、排除 ST/退市/北交所的依据)。""" + __tablename__ = "stock_basic" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + ts_code: Mapped[str] = mapped_column(String(12), unique=True, index=True) # 000001.SZ + symbol: Mapped[str] = mapped_column(String(10), index=True) # 000001 + name: Mapped[str] = mapped_column(String(32)) + area: Mapped[str | None] = mapped_column(String(32)) + industry: Mapped[str | None] = mapped_column(String(32)) + market: Mapped[str | None] = mapped_column(String(32)) # 主板/创业板/科创板/北交所 + exchange: Mapped[str] = mapped_column(String(8)) # SSE/SZSE/BSE + list_status: Mapped[str] = mapped_column(String(2), index=True) # L上市 D退市 P暂停 + list_date: Mapped[str] = mapped_column(String(8), default="") + delist_date: Mapped[str | None] = mapped_column(String(8)) + + +class MarketDaily(Base): + """全市场未复权日线(选股专用,与回测 candles(qfq) 隔离)。 + + 单位沿用 Tushare 原始:vol 手、amount 千元。 + """ + __tablename__ = "market_daily" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + trade_date: Mapped[datetime] = mapped_column(DateTime, index=True) + ts_code: Mapped[str] = mapped_column(String(12), index=True) + open: Mapped[float] = mapped_column(Float) + high: Mapped[float] = mapped_column(Float) + low: Mapped[float] = mapped_column(Float) + close: Mapped[float] = mapped_column(Float) + pre_close: Mapped[float] = mapped_column(Float) + change: Mapped[float | None] = mapped_column(Float) + pct_chg: Mapped[float | None] = mapped_column(Float) # 日涨跌幅 % + vol: Mapped[float] = mapped_column(Float) # 手 + amount: Mapped[float] = mapped_column(Float) # 千元 + + __table_args__ = ( + UniqueConstraint("ts_code", "trade_date", name="uq_mkt_code_date"), + ) + + +class DailySnapshot(Base): + """每日指标快照(daily_basic)。total_mv/circ_mv 单位万元(Tushare 原始),API 层换算亿元。""" + __tablename__ = "daily_snapshot" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + trade_date: Mapped[datetime] = mapped_column(DateTime, index=True) + ts_code: Mapped[str] = mapped_column(String(12), index=True) + close: Mapped[float | None] = mapped_column(Float) + turnover_rate: Mapped[float | None] = mapped_column(Float) # 换手率 % + turnover_rate_f: Mapped[float | None] = mapped_column(Float) # 自由流通换手率 % + volume_ratio: Mapped[float | None] = mapped_column(Float) # 量比 + pe: Mapped[float | None] = mapped_column(Float) + pe_ttm: Mapped[float | None] = mapped_column(Float) + pb: Mapped[float | None] = mapped_column(Float) + total_mv: Mapped[float | None] = mapped_column(Float) # 总市值(万元) + circ_mv: Mapped[float | None] = mapped_column(Float) # 流通市值(万元) + + __table_args__ = ( + UniqueConstraint("ts_code", "trade_date", name="uq_snap_code_date"), + ) + + +class TradeCalendar(Base): + """交易日历缓存(trade_cal 拉取一次宽范围后本地维护,低积分 token 限频 1 次/小时)。""" + __tablename__ = "trade_calendar" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + trade_date: Mapped[str] = mapped_column(String(8), unique=True, index=True) # YYYYMMDD diff --git a/backend/app/schemas.py b/backend/app/schemas.py index cf3fc72..12fdb25 100644 --- a/backend/app/schemas.py +++ b/backend/app/schemas.py @@ -5,6 +5,7 @@ from __future__ import annotations from datetime import datetime +from typing import Literal from pydantic import BaseModel, Field @@ -85,3 +86,84 @@ class SyncResponse(BaseModel): symbol: str bars: int source: str + + +# ---------- Screener(智能选股) ---------- +Op = Literal["gt", "ge", "lt", "le", "between"] + + +class IndicatorCondition(BaseModel): + """技术指标条件(在最近 lookback 个交易日窗口内判定)。 + + indicator 白名单见 screener/llm.py 的 SYSTEM_PROMPT(kdj_j / rsi / macd_dif…)。 + 设置 value_indicator 时为指标间比较(如 DIF > DEA、close < boll_lower),value 填 0 占位。 + """ + indicator: str + params: dict[str, float] = Field(default_factory=dict) # 如 {"n": 9, "m1": 3, "m2": 3} + op: Op + value: float + value2: float | None = None # between 上界 + value_indicator: str | None = None # 比较对象为另一指标(同白名单)时使用 + value_params: dict[str, float] = Field(default_factory=dict) # 比较对象指标参数(默认沿用 params/默认值) + lookback: int = 1 # 检查最近 N 个交易日 + match: Literal["all", "any"] = "all" # all=连续满足;any=任一满足 + + +class SnapshotCondition(BaseModel): + """每日快照条件(最新交易日截面)。市值单位亿元;换手率为百分数(5 表示 5%)。""" + field: str # total_mv|circ_mv|pe_ttm|pb|turnover_rate|close + op: Op + value: float + value2: float | None = None + + +class ScreenConditions(BaseModel): + indicator: list[IndicatorCondition] = Field(default_factory=list) + snapshot: list[SnapshotCondition] = Field(default_factory=list) + exclude_st: bool = True + exclude_delisted: bool = True + exclude_bj: bool = True # 排除北交所 + + +class ScreenerRunRequest(BaseModel): + text: str = Field(min_length=2, max_length=500) + # 直传条件则跳过 LLM 解析(预留给"微调再跑") + conditions: ScreenConditions | None = None + + +class ScreenerItemOut(BaseModel): + ts_code: str + name: str + close: float | None = None # 最新收盘价(元) + pct_chg: float | None = None # 日涨跌幅 % + total_mv: float | None = None # 总市值(亿元) + circ_mv: float | None = None # 流通市值(亿元) + pe_ttm: float | None = None + pb: float | None = None + turnover_rate: float | None = None + indicators: dict[str, float | None] = Field(default_factory=dict) # 引用到的指标最新值 + + +class ScreenerRunResponse(BaseModel): + conditions: ScreenConditions + trade_date: datetime | None # 数据基准交易日 + total: int # 命中总数(items 可能被截断) + items: list[ScreenerItemOut] + indicator_labels: dict[str, str] = Field(default_factory=dict) # "kdj_j" -> "KDJ J(9,3,3)" + + +class ScreenerSyncRequest(BaseModel): + days: int = Field(default=90, ge=10, le=250) # 同步最近 N 个交易日 + force: bool = False # True => 全量重拉(幂等) + + +class ScreenerSyncStatus(BaseModel): + running: bool + step: str | None = None # 进行中步骤文案 + total_days: int = 0 + done_days: int = 0 + error: str | None = None + ready: bool = False # 至少 1 个交易日数据可用于选股 + last_trade_date: datetime | None = None + last_synced_at: datetime | None = None + stats: dict[str, int] = Field(default_factory=dict) # stocks/daily_rows/snapshot_rows/dates diff --git a/backend/app/screener/__init__.py b/backend/app/screener/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/app/screener/engine.py b/backend/app/screener/engine.py new file mode 100644 index 0000000..7668cf5 --- /dev/null +++ b/backend/app/screener/engine.py @@ -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, + } diff --git a/backend/app/screener/llm.py b/backend/app/screener/llm.py new file mode 100644 index 0000000..a1a8c87 --- /dev/null +++ b/backend/app/screener/llm.py @@ -0,0 +1,151 @@ +"""LLM 条件解析器(DeepSeek,OpenAI 兼容 /chat/completions)。 + +把自然语言选股需求解析成 ScreenConditions(结构化 JSON)。 +- response_format=json_object + temperature=0.1 保证结构稳定 +- 解析/校验失败带错误重试 1 次 +- 上游错误信息透传给前端(ScreenerError) +""" +from __future__ import annotations + +import json +import re + +import httpx + +from ..config import settings +from ..schemas import ScreenConditions + +SYSTEM_PROMPT = """你是 A 股选股条件解析器。把用户的自然语言解析成一个 JSON 对象,只输出 JSON,不要任何解释、注释或代码块围栏。完全无法理解时输出 {"error": "原因"}。 + +输出结构: +{"indicator": [...], "snapshot": [...], "exclude_st": true, "exclude_delisted": true, "exclude_bj": true} +(indicator 与 snapshot 至少一个非空;用户没有提到的条件不要编造) + +【indicator 数组】技术指标条件,元素字段: +- "indicator": 指标名,白名单:kdj_k / kdj_d / kdj_j(KDJ 的 K/D/J 值)、rsi、macd_dif / macd_dea / macd_hist(MACD 的 DIF/DEA/柱)、ma(收盘价均线)、boll_upper / boll_mid / boll_lower(布林轨道)、close(收盘价)、pct_chg(日涨跌幅%) +- "params": 指标参数(可选),默认:KDJ {"n":9,"m1":3,"m2":3};RSI {"period":14};MACD {"fast":12,"slow":26,"signal":9};MA {"period":20};BOLL {"period":20,"std":2} +- "op": "gt" | "ge" | "lt" | "le" | "between" +- "value": 比较数值(between 时为下界),"value2": between 上界 +- "value_indicator": 可选。指标与指标比较时填另一指标名(同白名单),如 "DIF大于DEA" -> indicator=macd_dif, op=gt, value_indicator=macd_dea, value=0;"股价在布林带下轨之下" -> indicator=close, op=lt, value_indicator=boll_lower, value=0 +- "value_params": 可选。比较对象指标需要不同参数时指定,如 "MA5 上穿 MA20" -> indicator=ma, params={"period":5}, op=gt, value_indicator=ma, value_params={"period":20}, value=0;不填则比较对象沿用 params 中适用于它的参数或默认参数 +- "lookback": 检查最近 N 个交易日(默认 1) +- "match": "all"(窗口内每天满足,默认)或 "any"(窗口内任一天满足) + +【snapshot 数组】最新交易日截面条件,元素字段: +- "field": 白名单:total_mv(总市值)、circ_mv(流通市值)、pe_ttm(市盈率TTM)、pb(市净率)、turnover_rate(换手率)、close(最新价) +- "op"/"value"/"value2" 同上 + +【单位约定】市值条件统一用亿元(如"市值大于100亿,小于200亿"→ between 100~200);换手率用百分数值("换手率大于5%"→ value 5);价格类用元;pe/pb 用倍数。 + +【时间语义】"今天/今日"→ lookback=1;"这两天/最近N天/连续N日"→ lookback=N 且 match="all";"近N日内曾经/任一天"→ lookback=N 且 match="any"。只支持以最新交易日为终点的窗口,不要生成具体某一天的条件。 + +【排除规则】默认 exclude_st=true、exclude_delisted=true、exclude_bj=true;用户明确说"包含北交所/包含ST"时才把对应项设为 false。 + +【交叉类表述的近似】"MACD金叉/刚金叉"用 macd_dif gt macd_dea(value_indicator)+ 适当 lookback/match 近似;"跌破均线"用 close lt ma 近似;无法近似表达的复杂条件直接忽略,保留可表达的部分。 + +示例1: +输入:帮我找出这两天 KDJ 中的 J 小于 10,市值大于 100 亿,小于 200 亿的公司 +输出:{"indicator":[{"indicator":"kdj_j","params":{"n":9,"m1":3,"m2":3},"op":"lt","value":10,"lookback":2,"match":"all"}],"snapshot":[{"field":"total_mv","op":"between","value":100,"value2":200}],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true} + +示例2: +输入:RSI 低于 30,市盈率 TTM 小于 20 的公司 +输出:{"indicator":[{"indicator":"rsi","params":{"period":14},"op":"lt","value":30,"lookback":1,"match":"all"}],"snapshot":[{"field":"pe_ttm","op":"lt","value":20}],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true} + +示例3: +输入:近 5 天曾经 MACD 金叉(DIF 上穿 DEA),换手率大于 5%,流通市值小于 100 亿 +输出:{"indicator":[{"indicator":"macd_dif","params":{"fast":12,"slow":26,"signal":9},"op":"gt","value":0,"value_indicator":"macd_dea","lookback":5,"match":"any"}],"snapshot":[{"field":"turnover_rate","op":"gt","value":5},{"field":"circ_mv","op":"lt","value":100}],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true}""" + + +class ScreenerError(RuntimeError): + """选股链路可预期的业务错误(信息可直接透传给前端)。""" + + +def _endpoint(base_url: str) -> str: + """归一化 base_url -> 完整 chat/completions URL。 + + 兼容多种写法:DeepSeek/OpenAI 的 .../v1、智谱 GLM 的 .../v4、或直接给完整路径。 + """ + base = base_url.rstrip("/") + if base.endswith("/chat/completions"): + return base + if not re.search(r"/v\d+$", base): # 未带版本段则补 /v1(DeepSeek/OpenAI 惯例) + base += "/v1" + return f"{base}/chat/completions" + + +def _extract_json(content: str) -> dict: + """从 LLM 输出提取 JSON:剥代码围栏,或取首 { 到末 } 的子串。""" + text = content.strip() + if text.startswith("```"): + # 剥 ```json ... ``` 围栏 + text = text.split("```", 2)[1] + if text.startswith("json"): + text = text[4:] + text = text.strip() + if not text.startswith("{"): + start, end = text.find("{"), text.rfind("}") + if start < 0 or end <= start: + raise ValueError("输出中不含 JSON 对象") + text = text[start : end + 1] + obj = json.loads(text) + if not isinstance(obj, dict): + raise ValueError("JSON 不是对象") + return obj + + +def _build_messages(text: str, retry_error: str | None = None) -> list[dict]: + """构造 system + user 消息;retry 时附上一次解析错误要求修正。""" + user = f"解析以下选股需求:{text}" + if retry_error: + user += f"\n\n上一次输出无法通过校验,错误:{retry_error}。请修正后重新只输出 JSON。" + return [{"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user}] + + +async def _chat(messages: list[dict]) -> str: + """调 OpenAI 兼容接口,返回 assistant 文本。上游错误抛 ScreenerError。""" + body = { + "model": settings.llm_model, + "messages": messages, + "temperature": 0.1, + "response_format": {"type": "json_object"}, + "max_tokens": 2000, + } + headers = {"Authorization": f"Bearer {settings.llm_api_key}"} + async with httpx.AsyncClient(timeout=settings.llm_timeout) as client: + try: + r = await client.post(_endpoint(settings.llm_base_url), json=body, headers=headers) + except httpx.HTTPError as e: # 网络/超时 + raise ScreenerError(f"LLM 服务无法访问({settings.llm_base_url}): {e}") from e + if r.status_code >= 400: + detail = "" + try: + detail = r.json().get("error", {}).get("message", "") + except Exception: # noqa: BLE001 + detail = r.text[:200] + raise ScreenerError(f"LLM 接口错误 (HTTP {r.status_code}): {detail or '无详细信息'}") + try: + return r.json()["choices"][0]["message"]["content"] or "" + except (KeyError, IndexError, TypeError) as e: + raise ScreenerError(f"LLM 返回结构异常: {r.text[:200]}") from e + + +async def parse_conditions(text: str) -> ScreenConditions: + """主入口:自然语言 -> ScreenConditions。未配 key / 解析两次失败抛 ScreenerError。""" + if not settings.llm_api_key: + raise ScreenerError( + "未配置 LLM_API_KEY:请在 backend/.env 填入 DeepSeek API Key(platform.deepseek.com 获取)后重启后端" + ) + + retry_error: str | None = None + for _ in range(2): # 首次 + 失败重试 1 次 + content = await _chat(_build_messages(text, retry_error)) + try: + obj = _extract_json(content) + if "error" in obj and not obj.get("indicator") and not obj.get("snapshot"): + raise ScreenerError(f"AI 无法理解该选股需求:{obj['error']}") + return ScreenConditions.model_validate(obj) + except ScreenerError: + raise + except Exception as e: # noqa: BLE001 —— JSON/校验失败,带错误重试 + retry_error = str(e)[:300] + raise ScreenerError(f"AI 解析结果两次未通过校验,最后错误:{retry_error}") diff --git a/backend/app/screener/market_sync.py b/backend/app/screener/market_sync.py new file mode 100644 index 0000000..3604230 --- /dev/null +++ b/backend/app/screener/market_sync.py @@ -0,0 +1,299 @@ +"""全市场数据同步(选股专用,未复权;与回测 candles 表隔离)。 + +设计:trade_cal 取近 N 个交易日 -> 逐日 pro.daily(trade_date=...) / pro.daily_basic(trade_date=...) +一次返回全市场当日数据 -> 按 trade_date 删旧插新批量入库(幂等)。 +同步为进程内后台任务(MVP 不引入任务队列),前端轮询 /api/screener/sync/status。 + +daily 与 daily_basic 分步独立落库:daily_basic 积分不足时日线仍可用,错误写入状态不中断任务。 +""" +from __future__ import annotations + +import asyncio +import time +from datetime import datetime, timedelta + +from sqlalchemy import delete, func, insert, select +from sqlalchemy.ext.asyncio import AsyncSession + +from ..config import settings +from ..models import DailySnapshot, MarketDaily, StockBasic, TradeCalendar +from .llm import ScreenerError + +# 进程内单例任务状态(uvicorn --reload 单进程场景够用) +_sync_state: dict = { + "running": False, + "step": None, + "total_days": 0, + "done_days": 0, + "error": None, + "started_at": None, + "finished_at": None, +} +_sync_task: asyncio.Task | None = None +_sync_lock = asyncio.Lock() + +_BATCH = 5000 # executemany 分批行数 + +# Tushare 积分/权限不足的特征文案(daily_basic 常见门槛) +_PERM_MARKS = ("抱歉,您没有访问该项目权限", "积分", "权限") +# 频率超限特征(等待 62s 重试一次) +_RATE_MARKS = ("频率超限", "每分钟") + + +def _call_retry(fn, *args, **kwargs): + """同步调用 tushare 接口;「每分钟」级频率超限等 62s 重试一次(小时级限频直接抛)。""" + try: + return fn(*args, **kwargs) + except Exception as e: # noqa: BLE001 + msg = str(e) + if any(m in msg for m in _RATE_MARKS) and "小时" not in msg: + time.sleep(62) + return fn(*args, **kwargs) + raise + + +def _get_pro(): + """token 检查 + 返回 pro api 客户端(同步对象,调用需 to_thread 包裹)。""" + if not settings.tushare_token: + raise ScreenerError("未配置 TUSHARE_TOKEN,无法同步全市场数据(backend/.env)") + import tushare as ts + + ts.set_token(settings.tushare_token) + return ts.pro_api() + + +def _parse_d(s: str) -> datetime: + return datetime.strptime(str(s), "%Y%m%d") + + +def _fetch_calendar_sync(pro) -> list[str]: + """拉取宽范围交易日历(近 18 个月 + 未来 3 个月),返回 YYYYMMDD 列表。""" + time.sleep(settings.screener_sync_interval) + end = (datetime.now() + timedelta(days=90)).strftime("%Y%m%d") + start = (datetime.now() - timedelta(days=550)).strftime("%Y%m%d") + cal = _call_retry(pro.trade_cal, exchange="SSE", start_date=start, end_date=end, is_open="1") + return sorted(cal["cal_date"].tolist()) + + +async def _recent_trade_dates(session: AsyncSession, pro, days: int) -> list[str]: + """近 N 个交易日(YYYYMMDD,倒序)。日历本地缓存,仅在覆盖不到当天时刷新一次。 + + trade_cal 低积分版限频 1 次/小时:刷新被限频时沿用缓存(日历略旧无害—— + daily 对未生成日期返回空,同步会自然跳过)。 + """ + cached = (await session.execute(select(TradeCalendar.trade_date).order_by(TradeCalendar.trade_date.desc()))).scalars().all() + today = datetime.now().strftime("%Y%m%d") + have_today = bool(cached) and cached[0] >= today + + if not have_today: + try: + dates = await asyncio.to_thread(_fetch_calendar_sync, pro) + await session.execute(delete(TradeCalendar)) + await session.execute(insert(TradeCalendar), [{"trade_date": d} for d in dates]) + await session.commit() + cached = dates[::-1] + except Exception as e: # noqa: BLE001 —— 限频且无缓存时才致命 + if not cached: + raise ScreenerError(f"获取交易日历失败(且本地无缓存): {str(e)[:150]}") from e + _sync_state["step"] = "交易日历刷新受限,沿用本地缓存" + + recent = [d for d in cached if d <= today][:days] + if not recent: + raise ScreenerError("交易日历为空") + return recent + + +def _fetch_daily(pro, d: str) -> list[dict]: + """拉取某交易日全市场日线(未复权)。当日数据未生成(盘前/盘中)返回空。""" + time.sleep(settings.screener_sync_interval) + df = _call_retry(pro.daily, trade_date=d) + if df is None or df.empty: + return [] + rows = [] + for _, r in df.iterrows(): + rows.append({ + "trade_date": _parse_d(d), + "ts_code": r["ts_code"], + "open": float(r["open"]), "high": float(r["high"]), + "low": float(r["low"]), "close": float(r["close"]), + "pre_close": float(r["pre_close"]), + "change": None if r.get("change") != r.get("change") else float(r["change"]), + "pct_chg": None if r.get("pct_chg") != r.get("pct_chg") else float(r["pct_chg"]), + "vol": float(r["vol"]), # 手 + "amount": float(r["amount"]), # 千元 + }) + return rows + + +def _fetch_basic(pro, d: str) -> list[dict]: + """拉取某交易日每日指标快照(daily_basic,低积分版限频 1 次/分钟)。 + + 失败(积分不足等)时记录错误返回空,不拖垮日线同步。 + """ + time.sleep(settings.screener_sync_interval) + try: + df = _call_retry(pro.daily_basic, trade_date=d) + except Exception as e: # noqa: BLE001 + msg = str(e) + if any(m in msg for m in _PERM_MARKS): + _sync_state["error"] = ( + f"Tushare 无法获取每日指标(daily_basic):{msg[:150]}。" + "市值/市盈率等条件不可用;纯指标选股不受影响。" + ) + return [] + raise + if df is None or df.empty: + return [] + rows = [] + for _, r in df.iterrows(): + def _f(key: str) -> float | None: + v = r.get(key) + return None if v is None or v != v else float(v) + rows.append({ + "trade_date": _parse_d(d), + "ts_code": r["ts_code"], + "close": _f("close"), "turnover_rate": _f("turnover_rate"), + "turnover_rate_f": _f("turnover_rate_f"), "volume_ratio": _f("volume_ratio"), + "pe": _f("pe"), "pe_ttm": _f("pe_ttm"), "pb": _f("pb"), + "total_mv": _f("total_mv"), "circ_mv": _f("circ_mv"), # 万元 + }) + return rows + + +def _sync_stock_list_sync(pro) -> list[dict]: + """拉取在市股票列表。""" + time.sleep(settings.screener_sync_interval) + df = _call_retry(pro.stock_basic, exchange="", list_status="L", + fields="ts_code,symbol,name,area,industry,market,exchange,list_status,list_date,delist_date") + rows = [] + for _, r in df.iterrows(): + rows.append({ + "ts_code": r["ts_code"], "symbol": r["symbol"], "name": r["name"], + "area": r.get("area") or None, "industry": r.get("industry") or None, + "market": r.get("market") or None, "exchange": r["exchange"] or "", + "list_status": r["list_status"], "list_date": r.get("list_date") or "", + "delist_date": r.get("delist_date") or None, + }) + return rows + + +def _norm_date(v) -> str: + """把 DB 读出的 trade_date(可能是 datetime 或 str)归一为 YYYYMMDD。""" + if hasattr(v, "strftime"): + return v.strftime("%Y%m%d") + return str(v)[:10].replace("-", "") + + +async def _existing_dates(session: AsyncSession, model) -> set[str]: + """某表已落库的交易日集合(YYYYMMDD 字符串,便于比对)。""" + res = await session.execute(select(func.distinct(model.trade_date))) + return {_norm_date(r[0]) for r in res} + + +async def _replace_day(session: AsyncSession, model, rows: list[dict], d_str: str) -> None: + """按交易日删旧插新(幂等),executemany 分批。""" + d = _parse_d(d_str) + await session.execute(delete(model).where(model.trade_date == d)) + for i in range(0, len(rows), _BATCH): + await session.execute(insert(model), rows[i : i + _BATCH]) + await session.commit() + + +async def _run_sync(days: int, force: bool) -> None: + """后台任务主体:stock_basic -> 逐日日线 -> 最新交易日快照。异常写状态。 + + daily_basic 只拉最新交易日(快照条件仅作用于最新截面,且低积分 token 限频 1 次/分钟)。 + """ + from ..db import async_session # 延迟导入避免循环 + + try: + pro = await asyncio.to_thread(_get_pro) + + # 1) 股票列表(已有数据则跳过——stock_basic 低积分版限频 1 次/小时) + async with async_session() as session: + stocks_now = int(await session.scalar(select(func.count()).select_from(StockBasic)) or 0) + if stocks_now == 0 or force: + _sync_state["step"] = "正在同步股票列表" + try: + rows = await asyncio.to_thread(_sync_stock_list_sync, pro) + async with async_session() as session: + await session.execute(delete(StockBasic)) + for i in range(0, len(rows), _BATCH): + await session.execute(insert(StockBasic), rows[i : i + _BATCH]) + await session.commit() + except Exception as e: # noqa: BLE001 —— 受限时沿用现有列表继续 + if stocks_now > 0: + _sync_state["step"] = f"股票列表同步受限(沿用现有 {stocks_now} 只)" + else: + raise + + # 2) 逐交易日日线(增量;当日未生成则跳过) + async with async_session() as session: + dates = await _recent_trade_dates(session, pro, days) + have_daily = set() if force else await _existing_dates(session, MarketDaily) + todo = [d for d in dates if d not in have_daily] + _sync_state["total_days"] = len(todo) + _sync_state["done_days"] = 0 + + for d in todo: + _sync_state["step"] = f"正在同步 {d} 日线({_sync_state['done_days'] + 1}/{len(todo)})" + daily_rows = await asyncio.to_thread(_fetch_daily, pro, d) + if daily_rows: # 盘前/盘中等未生成数据的日期直接跳过 + async with async_session() as session: + await _replace_day(session, MarketDaily, daily_rows, d) + _sync_state["done_days"] += 1 + + # 3) 最新「有数据」交易日的快照(daily_basic,仅 1 次调用) + # 用 market_daily 实际最大交易日(今天的数据收盘后才生成,日历最新日会拉到空) + async with async_session() as session: + latest_dt = await session.scalar(select(func.max(MarketDaily.trade_date))) + latest = latest_dt.strftime("%Y%m%d") if latest_dt else None + if latest: + async with async_session() as session: + have_snap = force or latest not in await _existing_dates(session, DailySnapshot) + if have_snap: + _sync_state["step"] = f"正在同步 {latest} 每日指标" + basic_rows = await asyncio.to_thread(_fetch_basic, pro, latest) + if basic_rows: + async with async_session() as session: + await _replace_day(session, DailySnapshot, basic_rows, latest) + + _sync_state["step"] = "同步完成" + except Exception as e: # noqa: BLE001 + _sync_state["error"] = f"同步失败:{str(e)[:300]}" + _sync_state["step"] = "同步失败" + finally: + _sync_state["running"] = False + _sync_state["finished_at"] = datetime.now() + + +async def start_sync(session: AsyncSession, days: int, force: bool) -> dict: + """幂等启动后台同步任务;已在跑则直接返回当前状态。""" + global _sync_task + async with _sync_lock: + if _sync_state["running"] and _sync_task and not _sync_task.done(): + return dict(_sync_state) + _sync_state.update({ + "running": True, "step": "准备同步", "total_days": days, "done_days": 0, + "error": None, "started_at": datetime.now(), "finished_at": None, + }) + _sync_task = asyncio.create_task(_run_sync(days, force)) + return dict(_sync_state) + + +async def get_sync_status(session: AsyncSession) -> dict: + """合并任务状态 + DB 实况(最新交易日/行数/ready 标志),与 ScreenerSyncStatus DTO 对齐。""" + stocks = int(await session.scalar(select(func.count()).select_from(StockBasic)) or 0) + daily_rows = int(await session.scalar(select(func.count()).select_from(MarketDaily)) or 0) + snap_rows = int(await session.scalar(select(func.count()).select_from(DailySnapshot)) or 0) + last_daily = await session.scalar(select(func.max(MarketDaily.trade_date))) + n_dates = int(await session.scalar(select(func.count(func.distinct(MarketDaily.trade_date)))) or 0) + + status = dict(_sync_state) + status.update({ + "stats": {"stocks": stocks, "daily_rows": daily_rows, "snapshot_rows": snap_rows, "dates": n_dates}, + "last_trade_date": last_daily, + "last_synced_at": _sync_state.get("finished_at") or _sync_state.get("started_at"), + "ready": daily_rows > 0, + }) + return status diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 5294b63..ae8e314 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -14,6 +14,7 @@ dependencies = [ "numpy>=1.26", "pandas>=2.2", "tushare>=1.4", + "httpx>=0.28.1", ] [tool.uv] diff --git a/backend/smoke_test.py b/backend/smoke_test.py index 4c839ce..cd62161 100644 --- a/backend/smoke_test.py +++ b/backend/smoke_test.py @@ -52,3 +52,49 @@ with TestClient(app) as c: assert len(wd["candles"]) < len(d["candles"]) print("\n✅ 后端全链路自检通过(含周期聚合)") + + # ---------- 智能选股 ---------- + print("\n== 智能选股 ==") + + # 1) JSON 提取容错(不联网):代码围栏 / 多余文本 + from app.screener.llm import _extract_json + assert _extract_json('```json\n{"a": 1}\n```') == {"a": 1} + assert _extract_json('好的,结果如下:{"indicator": [], "snapshot": []} 谢谢')["indicator"] == [] + print("_extract_json 围栏/噪音容错 ✅") + + # 2) 未配置 LLM_API_KEY 时 /run 返回 503(确定性,不联网) + from app.config import settings as _s + r = c.post("/api/screener/run", json={"text": "这两天 KDJ 的 J 小于 10"}) + if not _s.llm_api_key: + assert r.status_code == 503, f"无 key 应 503,实际 {r.status_code}" + print("无 LLM_API_KEY -> 503 ✅") + else: + print("已配置 LLM_API_KEY,跳过 503 用例") + + # 3) 直传条件选股(不依赖 LLM;依赖已同步的全市场数据) + from app.models import MarketDaily # noqa: F401 + from sqlalchemy import select, func + from app.db import async_session + import asyncio + + async def _has_data() -> bool: + async with async_session() as session: + return (await session.scalar(select(func.count()).select_from(MarketDaily))) or 0 > 0 + + if asyncio.run(_has_data()): + r = c.post("/api/screener/run", json={ + "text": "测试直传", + "conditions": { + "indicator": [{"indicator": "kdj_j", "params": {"n": 9, "m1": 3, "m2": 3}, + "op": "lt", "value": 0, "lookback": 1, "match": "all"}], + "snapshot": [], + }, + }) + assert r.status_code == 200, f"直传选股失败 {r.status_code}: {r.text}" + d = r.json() + assert d["total"] >= 0 + print(f"KDJ J<0 选股 ✅ 命中 {d['total']} 只,基准日 {(d['trade_date'] or '')[:10]}") + else: + print("(未同步全市场数据,跳过直传选股用例;运行 POST /api/screener/sync 后再试)") + + print("\n✅ 智能选股自检通过") diff --git a/backend/uv.lock b/backend/uv.lock index 184584d..8983a59 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -286,6 +286,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/04/4b/29cac41a4d98d144bf5f6d33995617b185d14b22401f75ca86f384e87ff1/h11-0.16.0-py3-none-any.whl", hash = "sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86", size = 37515, upload-time = "2025-04-24T03:35:24.344Z" }, ] +[[package]] +name = "httpcore" +version = "1.0.9" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "certifi" }, + { name = "h11" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/06/94/82699a10bca87a5556c9c59b5963f2d039dbd239f25bc2a63907a05a14cb/httpcore-1.0.9.tar.gz", hash = "sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8", size = 85484, upload-time = "2025-04-24T22:06:22.219Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/7e/f5/f66802a942d491edb555dd61e3a9961140fd64c90bce1eafd741609d334d/httpcore-1.0.9-py3-none-any.whl", hash = "sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55", size = 78784, upload-time = "2025-04-24T22:06:20.566Z" }, +] + [[package]] name = "httptools" version = "0.8.0" @@ -322,6 +335,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/48/63/b906c01e53f50d432c0defe43ce52764a111dc1bdd028bafbeb54dcfd008/httptools-0.8.0-cp314-cp314t-win_amd64.whl", hash = "sha256:384c17174464c8e873398b7af24f0b1f44d992c820328413951a625323155d77", size = 108209, upload-time = "2026-05-25T22:17:39.473Z" }, ] +[[package]] +name = "httpx" +version = "0.28.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "certifi" }, + { name = "httpcore" }, + { name = "idna" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/b1/df/48c586a5fe32a0f01324ee087459e112ebb7224f646c0b5023f5e79e9956/httpx-0.28.1.tar.gz", hash = "sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc", size = 141406, upload-time = "2024-12-06T15:37:23.222Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2a/39/e50c7c3a983047577ee07d2a9e53faf5a69493943ec3f6a384bdc792deb2/httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad", size = 73517, upload-time = "2024-12-06T15:37:21.509Z" }, +] + [[package]] name = "idna" version = "3.18" @@ -827,6 +855,7 @@ dependencies = [ { name = "aiosqlite" }, { name = "asyncpg" }, { name = "fastapi" }, + { name = "httpx" }, { name = "numpy" }, { name = "pandas" }, { name = "pydantic" }, @@ -841,6 +870,7 @@ requires-dist = [ { name = "aiosqlite", specifier = ">=0.20" }, { name = "asyncpg", specifier = ">=0.29" }, { name = "fastapi", specifier = ">=0.115" }, + { name = "httpx", specifier = ">=0.28.1" }, { name = "numpy", specifier = ">=1.26" }, { name = "pandas", specifier = ">=2.2" }, { name = "pydantic", specifier = ">=2.7" }, diff --git a/frontend/package.json b/frontend/package.json index d07d707..e0f88fa 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -16,7 +16,8 @@ "pinia": "^2.3.0", "primeicons": "^7.0.0", "primevue": "^5.0.0", - "vue": "^3.5.0" + "vue": "^3.5.0", + "vue-router": "^4.6.4" }, "devDependencies": { "@types/node": "^22.0.0", diff --git a/frontend/pnpm-lock.yaml b/frontend/pnpm-lock.yaml index 7574505..5d1a1e4 100644 --- a/frontend/pnpm-lock.yaml +++ b/frontend/pnpm-lock.yaml @@ -29,6 +29,9 @@ importers: vue: specifier: ^3.5.0 version: 3.5.41(typescript@5.6.2) + vue-router: + specifier: ^4.6.4 + version: 4.6.4(vue@3.5.41(typescript@5.6.2)) devDependencies: '@types/node': specifier: ^22.0.0 @@ -457,6 +460,9 @@ packages: '@vue/devtools-api@6.6.3': resolution: {integrity: sha512-0MiMsFma/HqA6g3KLKn+AGpL1kgKhFWszC9U29NfpWK5LE7bjeXxySWJrOJ77hBz+TBrBQ7o4QJqbPbqbs8rJw==} + '@vue/devtools-api@6.6.4': + resolution: {integrity: sha512-sGhTPMuXqZ1rVOk32RylztWkfXTRhuS7vgAKv0zjqk8gbsHkJ7xfFf+jbySxt7tWObEJwyKaHMikV/WGDiQm8g==} + '@vue/language-core@2.1.0': resolution: {integrity: sha512-S7uUdQXn4aA7QloQeAIhKDkP243nCEktQu1Kxur1fJFfedf+izBTb9bAR2tHe0V7xhFYhIwkNCtQl2AlKFaW3w==} peerDependencies: @@ -662,6 +668,11 @@ packages: '@vue/composition-api': optional: true + vue-router@4.6.4: + resolution: {integrity: sha512-Hz9q5sa33Yhduglwz6g9skT8OBPii+4bFn88w6J+J4MfEo4KRRpmiNG/hHHkdbRFlLBOqxN8y8gf2Fb0MTUgVg==} + peerDependencies: + vue: ^3.5.0 + vue-tsc@2.1.0: resolution: {integrity: sha512-N6kelLgHPLiDzkeWfsfrfR79POitZ59M/2HaRWk4BW0ED40FlyQaVhGUvjg8VRPPJWic1M974cRTGIBR6zfgwQ==} hasBin: true @@ -967,6 +978,8 @@ snapshots: '@vue/devtools-api@6.6.3': {} + '@vue/devtools-api@6.6.4': {} + '@vue/language-core@2.1.0(typescript@5.6.2)': dependencies: '@volar/language-core': 2.4.28 @@ -1182,6 +1195,11 @@ snapshots: dependencies: vue: 3.5.41(typescript@5.6.2) + vue-router@4.6.4(vue@3.5.41(typescript@5.6.2)): + dependencies: + '@vue/devtools-api': 6.6.4 + vue: 3.5.41(typescript@5.6.2) + vue-tsc@2.1.0(typescript@5.6.2): dependencies: '@volar/typescript': 2.4.28 diff --git a/frontend/src/App.vue b/frontend/src/App.vue index 54094d7..c06dfef 100644 --- a/frontend/src/App.vue +++ b/frontend/src/App.vue @@ -1,13 +1,18 @@ diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index 35f6202..67f13b8 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -1,4 +1,13 @@ -import type { BacktestRequest, BacktestResponse, SyncRequest, SyncResponse } from './types'; +import type { + BacktestRequest, + BacktestResponse, + ScreenerRunRequest, + ScreenerRunResponse, + ScreenerSyncRequest, + ScreenerSyncStatus, + SyncRequest, + SyncResponse, +} from './types'; // dev 用 Vite 代理(/api -> :8000);生产构建设 VITE_API_BASE 指向后端地址。 const BASE = import.meta.env.VITE_API_BASE ?? ''; @@ -27,3 +36,47 @@ export async function syncData(req: SyncRequest): Promise { return (await res.json()) as SyncResponse; } +// ---------- 智能选股 ---------- +export async function runScreener(req: ScreenerRunRequest): Promise { + const res = await fetch(`${BASE}/api/screener/run`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(req), + }); + if (!res.ok) { + // 后端 detail 字段携带可读中文原因(503 未配 key / 409 未同步 / 502 LLM 错) + let detail = ''; + try { + detail = (await res.json())?.detail ?? ''; + } catch { + detail = await res.text(); + } + throw new Error(detail || `选股失败 (HTTP ${res.status})`); + } + return (await res.json()) as ScreenerRunResponse; +} + +export async function startScreenerSync(req: ScreenerSyncRequest = {}): Promise { + const res = await fetch(`${BASE}/api/screener/sync`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(req), + }); + if (!res.ok) { + let detail = ''; + try { + detail = (await res.json())?.detail ?? ''; + } catch { + detail = await res.text(); + } + throw new Error(detail || `启动同步失败 (HTTP ${res.status})`); + } + return (await res.json()) as ScreenerSyncStatus; +} + +export async function getScreenerSyncStatus(): Promise { + const res = await fetch(`${BASE}/api/screener/sync/status`); + if (!res.ok) throw new Error(`获取同步状态失败 (HTTP ${res.status})`); + return (await res.json()) as ScreenerSyncStatus; +} + diff --git a/frontend/src/api/types.ts b/frontend/src/api/types.ts index 03bb56d..4943a3e 100644 --- a/frontend/src/api/types.ts +++ b/frontend/src/api/types.ts @@ -74,3 +74,76 @@ export interface SyncResponse { bars: number; source: string; } + +// ---------- 智能选股(镜像 app/schemas.py) ---------- +export type Op = 'gt' | 'ge' | 'lt' | 'le' | 'between'; + +export interface IndicatorCondition { + indicator: string; + params?: Record; + op: Op; + value: number; + value2?: number | null; + value_indicator?: string | null; + value_params?: Record; + lookback?: number; + match?: 'all' | 'any'; +} + +export interface SnapshotCondition { + field: string; // total_mv | circ_mv | pe_ttm | pb | turnover_rate | close + op: Op; + value: number; + value2?: number | null; +} + +export interface ScreenConditions { + indicator: IndicatorCondition[]; + snapshot: SnapshotCondition[]; + exclude_st: boolean; + exclude_delisted: boolean; + exclude_bj: boolean; +} + +export interface ScreenerRunRequest { + text: string; + conditions?: ScreenConditions | null; // 直传则跳过 LLM 解析 +} + +export interface ScreenerItemOut { + ts_code: string; + name: string; + close: number | null; + pct_chg: number | null; + total_mv: number | null; // 亿元 + circ_mv: number | null; + pe_ttm: number | null; + pb: number | null; + turnover_rate: number | null; + indicators: Record; +} + +export interface ScreenerRunResponse { + conditions: ScreenConditions; + trade_date: string | null; + total: number; + items: ScreenerItemOut[]; + indicator_labels: Record; +} + +export interface ScreenerSyncRequest { + days?: number; + force?: boolean; +} + +export interface ScreenerSyncStatus { + running: boolean; + step?: string | null; + total_days: number; + done_days: number; + error?: string | null; + ready: boolean; + last_trade_date?: string | null; + last_synced_at?: string | null; + stats: { stocks: number; daily_rows: number; snapshot_rows: number; dates: number }; +} diff --git a/frontend/src/components/ConditionChips.vue b/frontend/src/components/ConditionChips.vue new file mode 100644 index 0000000..4253c9c --- /dev/null +++ b/frontend/src/components/ConditionChips.vue @@ -0,0 +1,51 @@ + + + diff --git a/frontend/src/components/ScreenerForm.vue b/frontend/src/components/ScreenerForm.vue new file mode 100644 index 0000000..8fefbac --- /dev/null +++ b/frontend/src/components/ScreenerForm.vue @@ -0,0 +1,50 @@ + + +