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

@@ -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})

View File

@@ -15,6 +15,17 @@ class Settings(BaseSettings):
data_adjust: str = "qfq" # 复权qfq 前复权 / hfq 后复权 / "" 不复权
data_default_start: str = "20200101" # 默认拉取起点(约近 5 年)
# ---- LLM智能选股的自然语言解析DeepSeekOpenAI 兼容协议,可换任意兼容网关)----
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 调整)

View File

@@ -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

View File

@@ -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_PROMPTkdj_j / rsi / macd_dif…
设置 value_indicator 时为指标间比较(如 DIF > DEA、close < boll_lowervalue 填 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

View File

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

151
backend/app/screener/llm.py Normal file
View File

@@ -0,0 +1,151 @@
"""LLM 条件解析器DeepSeekOpenAI 兼容 /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_jKDJ 的 K/D/J 值、rsi、macd_dif / macd_dea / macd_histMACD 的 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_deavalue_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): # 未带版本段则补 /v1DeepSeek/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 Keyplatform.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}")

View File

@@ -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