"""事件回测引擎:入场条件命中 -> 次日买入 -> 持有 N 日 -> 全市场汇总统计。 数据口径: - 行情底座是 candles(TDX 全量导入,不复权),全历史可用; - 指标计算用不复权价(与选股/看盘口径一致:J<10、RSI<30 等阈值均为归一化或惯例值); - 收益率用 adj_factor 校正(ret = 出场价×f出 / 入场价×f入 - 1),消除除权除息失真; 因子缺失的股退化为不复权收益(新股/缺因子,样本中占少数)。 信号语义:与选股引擎一致——每条条件在信号日 d 为终点、lookback 窗口内 match=all(连续满足)/any(曾经满足),多条件之间取 AND。 """ from __future__ import annotations from datetime import date, datetime, timedelta import numpy as np import pandas as pd from sqlalchemy import and_, func, not_, or_, select from sqlalchemy.ext.asyncio import AsyncSession from ..models import AdjFactor, Candle, StockBasic from ..schemas import EventBacktestSpec from ..screener.engine import ( FAMILIES, _family_of, _op_mask, _params_for, _resolve_params, _series_for, ) # 指标配热缓冲 bar 数(MACD 等 EMA 类指标需要较长窗口才收敛) BUFFER_BARS = 80 # 每批查询的股票数(全市场分块拉取,避免单条 SQL 过大) BATCH_SIZE = 800 # 单次回测允许的最大样本数(超过则仅按日期取最近的,防内存失控) MAX_TRADES = 200_000 class EventEngineError(RuntimeError): """事件回测可预期的业务错误(信息透传前端)。""" def _signal_mask(g: pd.DataFrame, spec: EventBacktestSpec, cache: dict) -> pd.Series: """单股全序列信号掩码:各条件(rolling lookback)AND。""" total = pd.Series(True, index=g.index) for cond in spec.entry.indicator: fam = _family_of(cond.indicator) if len(g) < FAMILIES[fam].min_bars: return pd.Series(False, index=g.index) s = _series_for(g, cond.indicator, _params_for(cond.indicator, cond.params), cache) if s is None: return pd.Series(False, index=g.index) if cond.value_indicator: target = _series_for(g, cond.value_indicator, _resolve_params(cond, cond.value_indicator), cache) if target is None: return pd.Series(False, index=g.index) else: target = pd.Series(cond.value, index=s.index) m = _op_mask(s, target, cond).astype(int) n = max(1, cond.lookback) if n > 1: rolled = m.rolling(n, min_periods=n).sum() m = (rolled == n) if cond.match == "all" else (rolled > 0) else: m = m.astype(bool) total = total & m.fillna(False).astype(bool) return total def _entry_exit_indices(sig_idx: int, spec: EventBacktestSpec, n: int) -> tuple[int, int] | None: """信号日索引 -> (入场索引, 出场索引)。前视/越界返回 None。""" entry_i = sig_idx + 1 # 信号收盘后才动手:一律次日 exit_i = entry_i + spec.holding_days if exit_i >= n: return None return entry_i, exit_i def _price_at(row: pd.Series, timing: str) -> float: return float(row["open"] if timing == "open" else row["close"]) def _stats_block(trades: list[dict]) -> dict: """样本集合 -> 汇总统计(空样本给零值)。""" if not trades: return { "samples": 0, "stocks": 0, "mean_pct": 0.0, "median_pct": 0.0, "win_rate": 0.0, "std_pct": 0.0, "p10_pct": 0.0, "p25_pct": 0.0, "p75_pct": 0.0, "p90_pct": 0.0, "max_pct": 0.0, "min_pct": 0.0, "by_year": [], } rets = np.array([t["ret_pct"] for t in trades], dtype=float) by_year: list[dict] = [] df = pd.DataFrame(trades) for year, grp in df.groupby(df["entry_date"].dt.year): r = grp["ret_pct"].to_numpy() by_year.append({ "year": int(year), "samples": int(len(r)), "mean_pct": round(float(r.mean()), 3), "median_pct": round(float(np.median(r)), 3), "win_rate": round(float((r > 0).mean() * 100), 2), }) by_year.sort(key=lambda x: x["year"]) return { "samples": int(len(rets)), "stocks": int(df["ts_code"].nunique()), "mean_pct": round(float(rets.mean()), 3), "median_pct": round(float(np.median(rets)), 3), "win_rate": round(float((rets > 0).mean() * 100), 2), "std_pct": round(float(rets.std(ddof=1)) if len(rets) > 1 else 0.0, 3), "p10_pct": round(float(np.percentile(rets, 10)), 3), "p25_pct": round(float(np.percentile(rets, 25)), 3), "p75_pct": round(float(np.percentile(rets, 75)), 3), "p90_pct": round(float(np.percentile(rets, 90)), 3), "max_pct": round(float(rets.max()), 3), "min_pct": round(float(rets.min()), 3), "by_year": by_year, } async def run_event_backtest( session: AsyncSession, spec: EventBacktestSpec, ts_code: str | None = None, start: date | None = None, end: date | None = None, ) -> dict: """主入口:返回 {spec, universe, start, end, stats, trades(sample), total}。""" entry = spec.entry if not entry.indicator: raise EventEngineError("入场条件必须包含技术指标条件(如 J<10、RSI<30)") # 时间窗:默认最近一年;end 以 candles 最大日期为准 end_dt = end if end_dt is None: end_dt = (await session.scalar(select(func.max(Candle.ts)))) or date.today() if isinstance(end_dt, datetime): end_dt = end_dt.date() start_dt = start or (end_dt - timedelta(days=365)) if start_dt >= end_dt: raise EventEngineError("回测起始日期必须早于结束日期") needed = _max_needed_bars_safe(entry) + BUFFER_BARS buffer_start = start_dt - timedelta(days=int(needed * 1.7)) # 交易日->日历日近似 # 股票池:ts_code+symbol 映射(candles 按 symbol 存) name_map: dict[str, str] = {} if ts_code: rows = (await session.execute( select(StockBasic.ts_code, StockBasic.symbol, StockBasic.name) .where(StockBasic.ts_code == ts_code) )).all() if not rows: raise EventEngineError(f"未知股票代码: {ts_code}") universe = [(r[0], r[1]) for r in rows] name_map = {r[0]: r[2] for r in rows} else: stmt = select(StockBasic.ts_code, StockBasic.symbol, StockBasic.name).where( StockBasic.list_status == "L" ) if entry.exclude_st: stmt = stmt.where(not_(or_(StockBasic.name.like("%ST%"), StockBasic.name.like("%退%")))) if entry.exclude_bj: stmt = stmt.where(not_(StockBasic.ts_code.like("%.BJ"))) rows = (await session.execute(stmt)).all() universe = [(r[0], r[1]) for r in rows] name_map = {r[0]: r[2] for r in rows} start_ts = datetime(start_dt.year, start_dt.month, start_dt.day) end_ts = datetime(end_dt.year, end_dt.month, end_dt.day, 23, 59, 59) buffer_ts = datetime(buffer_start.year, buffer_start.month, buffer_start.day) trades: list[dict] = [] for i in range(0, len(universe), BATCH_SIZE): batch = universe[i : i + BATCH_SIZE] symbols = [sym for _, sym in batch] code_by_symbol = {sym: code for code, sym in batch} candle_rows = (await session.execute( select(Candle.symbol, Candle.ts, Candle.open, Candle.high, Candle.low, Candle.close) .where(and_(Candle.timeframe == "1d", Candle.symbol.in_(symbols), Candle.ts >= buffer_ts, Candle.ts <= end_ts)) .order_by(Candle.symbol, Candle.ts) )).all() if not candle_rows: continue codes = {code_by_symbol[s] for s in symbols} adj_rows = (await session.execute( select(AdjFactor.ts_code, AdjFactor.trade_date, AdjFactor.adj_factor) .where(and_(AdjFactor.ts_code.in_(codes), AdjFactor.trade_date >= buffer_ts, AdjFactor.trade_date <= end_ts)) )).all() f_map = {(r[0], r[1].date()): float(r[2]) for r in adj_rows if r[2]} bars = pd.DataFrame( candle_rows, columns=["symbol", "ts", "open", "high", "low", "close"] ) for symbol, g in bars.groupby("symbol", sort=False): if len(g) < 30: continue g = g.reset_index(drop=True) ts_code_l = code_by_symbol[symbol] cache: dict = {"_families": set()} mask = _signal_mask(g, spec, cache) if not mask.any(): continue for sig_i in np.flatnonzero(mask.to_numpy()): ts_sig = g.at[sig_i, "ts"] # 信号必须落在回测窗口内(buffer 区只用于指标配热) if ts_sig < start_ts: continue ie = _entry_exit_indices(int(sig_i), spec, len(g)) if ie is None: continue entry_i, exit_i = ie e_row, x_row = g.iloc[entry_i], g.iloc[exit_i] e_price = _price_at(e_row, "open" if spec.entry_timing == "next_open" else "close") x_price = _price_at(x_row, "open" if spec.exit_timing == "open" else "close") if not e_price or not x_price: continue f_in = f_map.get((ts_code_l, e_row["ts"].date()), 1.0) f_out = f_map.get((ts_code_l, x_row["ts"].date()), 1.0) ret_pct = (x_price * f_out) / (e_price * f_in) * 100 - 100 trades.append({ "ts_code": ts_code_l, "name": name_map.get(ts_code_l), "entry_date": e_row["ts"], "entry_price": round(e_price, 3), "exit_date": x_row["ts"], "exit_price": round(x_price, 3), "ret_pct": round(float(ret_pct), 3), }) if len(trades) >= MAX_TRADES: break if len(trades) >= MAX_TRADES: break if len(trades) >= MAX_TRADES: break stats = _stats_block(trades) # 明细样本:最好 100 + 最差 100(其余统计已覆盖) trades_sorted = sorted(trades, key=lambda t: t["ret_pct"], reverse=True) sample = trades_sorted[:100] + (trades_sorted[-100:] if len(trades_sorted) > 100 else []) return { "spec": spec, "universe": ts_code or "all", "start": start_ts, "end": end_ts, "stats": stats, "trades": sample, "total": stats["samples"], } # ---------- 小工具 ---------- def _max_needed_bars_safe(conds) -> int: """指标配热所需最大 bar 数(同 screener.engine._max_needed_bars)。""" 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