看股功能更新
This commit is contained in:
@@ -1,6 +1,6 @@
|
||||
"""选股执行引擎:SQL 快照预筛缩小范围 -> 逐股指标计算过滤。
|
||||
|
||||
性能:预筛在 SQLite 索引上完成(毫秒级);指标阶段候选集通常数百~数千只 × ~90 根 bar,
|
||||
性能:预筛在数据库索引上完成(毫秒级);指标阶段候选集通常数百~数千只 × ~90 根 bar,
|
||||
pandas 逐股计算(复用 app/indicators,指标按 (族, 参数) 去重计算),秒级完成。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
@@ -14,7 +14,8 @@ 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 ..data.symbols import plain_code
|
||||
from ..models import Candle, DailySnapshot, StockBasic
|
||||
from ..schemas import IndicatorCondition, ScreenConditions, SnapshotCondition
|
||||
|
||||
# 快照字段 -> (DB 列, 中文标签, LLM 值 -> DB 值的换算乘数)
|
||||
@@ -25,7 +26,7 @@ SNAPSHOT_FIELDS: dict[str, tuple[str, str, float]] = {
|
||||
"pe_ttm": ("pe_ttm", "市盈率TTM", 1.0),
|
||||
"pb": ("pb", "市净率", 1.0),
|
||||
"turnover_rate": ("turnover_rate", "换手率%", 1.0),
|
||||
"close": ("close", "最新价", 1.0), # 实际取 market_daily.close,无需快照表
|
||||
"close": ("close", "最新价", 1.0), # 实际取 candles 最新 bar,无需快照表
|
||||
}
|
||||
|
||||
|
||||
@@ -198,7 +199,7 @@ def _snapshot_clause(cond: SnapshotCondition):
|
||||
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)
|
||||
col = Candle.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":
|
||||
@@ -219,17 +220,22 @@ _CAND_COLS = ["ts_code", "name", "close", "pct_chg",
|
||||
|
||||
|
||||
async def _prefilter(session: AsyncSession, conds: ScreenConditions, target_date: datetime) -> pd.DataFrame:
|
||||
"""最新交易日截面预筛:快照条件 + 排除项 + 名称/收盘价。返回候选 DataFrame。"""
|
||||
"""最新交易日截面预筛:candles(全量不复权底座) + stock_basic + daily_snapshot。
|
||||
|
||||
只取 target_date 当日有交易的股票(停牌股无当日 bar,自然排除,与原先 market_daily
|
||||
的 trade_date == target_date 行为一致)。pct_chg 用前一交易日收盘价计算。
|
||||
"""
|
||||
snap_date = await session.scalar(select(func.max(DailySnapshot.trade_date))) or target_date
|
||||
stmt = (
|
||||
select(
|
||||
MarketDaily.ts_code, StockBasic.name, MarketDaily.close, MarketDaily.pct_chg,
|
||||
Candle.symbol, StockBasic.ts_code, StockBasic.name, Candle.close,
|
||||
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)
|
||||
.join(StockBasic, StockBasic.symbol == Candle.symbol)
|
||||
.outerjoin(DailySnapshot, and_(DailySnapshot.ts_code == StockBasic.ts_code,
|
||||
DailySnapshot.trade_date == snap_date))
|
||||
.where(Candle.timeframe == "1d", Candle.ts == target_date)
|
||||
)
|
||||
|
||||
if conds.exclude_delisted:
|
||||
@@ -237,13 +243,39 @@ async def _prefilter(session: AsyncSession, conds: ScreenConditions, target_date
|
||||
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")))
|
||||
stmt = stmt.where(not_(StockBasic.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)
|
||||
df = pd.DataFrame(rows, columns=["symbol", "ts_code", "name", "close",
|
||||
"total_mv", "circ_mv", "pe_ttm", "pb", "turnover_rate"])
|
||||
if df.empty:
|
||||
return df[_CAND_COLS]
|
||||
|
||||
# pct_chg:candles 无现成涨跌幅列,用前一交易日的收盘价计算
|
||||
prev_dt = await session.scalar(
|
||||
select(func.max(Candle.ts)).where(Candle.timeframe == "1d", Candle.ts < target_date)
|
||||
)
|
||||
prev_map: dict[str, float] = {}
|
||||
if prev_dt is not None:
|
||||
pr = await session.execute(
|
||||
select(Candle.symbol, Candle.close).where(
|
||||
Candle.timeframe == "1d", Candle.ts == prev_dt,
|
||||
Candle.symbol.in_(df["symbol"].tolist()),
|
||||
)
|
||||
)
|
||||
prev_map = {r.symbol: r.close for r in pr}
|
||||
prev = df["symbol"].map(prev_map)
|
||||
|
||||
def _pct(c, p) -> float | None:
|
||||
if p is None or p != p or float(p) == 0:
|
||||
return None
|
||||
return (float(c) / float(p) - 1) * 100
|
||||
|
||||
df["pct_chg"] = [_pct(c, p) for c, p in zip(df["close"], prev)]
|
||||
return df[_CAND_COLS]
|
||||
|
||||
|
||||
# ---------- 主流程 ----------
|
||||
@@ -260,18 +292,30 @@ def _max_needed_bars(conds: ScreenConditions) -> int:
|
||||
|
||||
async def _load_bars(session: AsyncSession, ts_codes: list[str],
|
||||
target_date: datetime, min_date: datetime) -> pd.DataFrame:
|
||||
"""载入候选股的 K 线窗口。候选 <= 2000 用 IN 精确圈定;否则拉全窗口再 pandas 过滤。"""
|
||||
"""载入候选股的 K 线窗口(candles 全量不复权底座)。
|
||||
|
||||
候选 <= 2000 用 IN 精确圈定;否则拉全窗口再 pandas 过滤。
|
||||
ts_code 按候选集映射回 symbol 查询,pct_chg 用每股收盘价环比计算。
|
||||
"""
|
||||
symbols = [plain_code(t) for t in ts_codes]
|
||||
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)))
|
||||
Candle.symbol, Candle.ts, Candle.open, Candle.high, Candle.low, Candle.close,
|
||||
).where(Candle.timeframe == "1d", Candle.ts >= min_date, Candle.ts <= target_date)
|
||||
if len(symbols) <= 2000:
|
||||
stmt = stmt.where(Candle.symbol.in_(set(symbols)))
|
||||
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)
|
||||
df = pd.DataFrame(rows, columns=["symbol", "trade_date", "open", "high", "low", "close"])
|
||||
if not df.empty and len(symbols) > 2000:
|
||||
df = df[df["symbol"].isin(set(symbols))]
|
||||
df = df.sort_values(["symbol", "trade_date"]).reset_index(drop=True)
|
||||
if df.empty:
|
||||
df["ts_code"] = pd.Series(dtype=object)
|
||||
df["pct_chg"] = pd.Series(dtype=object)
|
||||
else:
|
||||
df["pct_chg"] = df.groupby("symbol")["close"].pct_change() * 100
|
||||
ts_map = {plain_code(t): t for t in ts_codes}
|
||||
df["ts_code"] = df["symbol"].map(ts_map)
|
||||
return df[["ts_code", "trade_date", "open", "high", "low", "close", "pct_chg"]]
|
||||
|
||||
|
||||
def _f(v) -> float | None:
|
||||
@@ -303,7 +347,9 @@ async def run_screen(session: AsyncSession, conds: ScreenConditions, limit: int)
|
||||
if not conds.indicator and not conds.snapshot:
|
||||
raise ValueError("筛选条件为空")
|
||||
|
||||
target_date = await session.scalar(select(func.max(MarketDaily.trade_date)))
|
||||
target_date = await session.scalar(
|
||||
select(func.max(Candle.ts)).where(Candle.timeframe == "1d")
|
||||
)
|
||||
if target_date is None:
|
||||
raise DataNotReadyError("全市场数据未同步:请先在选股页点击「同步市场数据」")
|
||||
|
||||
@@ -330,10 +376,11 @@ async def run_screen(session: AsyncSession, conds: ScreenConditions, limit: int)
|
||||
# 纯快照条件:预筛结果即命中
|
||||
items = [_item_from_row(row) | {"indicators": {}} for _, row in cand.iterrows()]
|
||||
else:
|
||||
# 圈定 K 线窗口:按已同步交易日序列回溯 needed 根
|
||||
# 圈定 K 线窗口:按 candles 全量交易日序列回溯 needed 根(不再受同步窗口限制)
|
||||
need = _max_needed_bars(conds)
|
||||
dates_res = await session.execute(
|
||||
select(MarketDaily.trade_date).distinct().order_by(MarketDaily.trade_date.desc()).limit(need)
|
||||
select(Candle.ts).where(Candle.timeframe == "1d").distinct()
|
||||
.order_by(Candle.ts.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)
|
||||
|
||||
Reference in New Issue
Block a user