diff --git a/src/easy_tdx/web/backtest_schemas.py b/src/easy_tdx/web/backtest_schemas.py index ebcf853..2f17978 100644 --- a/src/easy_tdx/web/backtest_schemas.py +++ b/src/easy_tdx/web/backtest_schemas.py @@ -25,6 +25,10 @@ __all__ = [ "SavedStrategyListResponse", "MultiStrategyItem", "MultiStrategyBacktestRequest", + "SignalScanRequest", + "SignalScanRecentSignal", + "SignalScanRow", + "SignalScanResult", "serialize_result", ] @@ -372,6 +376,59 @@ class MultiStrategyBacktestRequest(BaseModel): execution: Literal["next_open", "next_close"] = Field(default="next_open") +# ── 信号雷达(一键扫描已保存策略的最近买卖信号)──────────────────────────────── + + +class SignalScanRequest(BaseModel): + """信号扫描请求:扫描策略库全部已保存策略,只看最近 N 根 K 线内的信号。""" + + window_bars: int = Field( + default=5, ge=1, le=30, description="检查最近 N 根 K 线内的信号(日线即 N 个交易日)" + ) + + +class SignalScanRecentSignal(BaseModel): + """窗口内单根 K 线的信号。""" + + date: str + direction: Literal["BUY", "SELL"] + + +class SignalScanRow(BaseModel): + """扫描结果单行:一个"策略×标的"子任务的信号摘要。 + + single 策略 1 行;portfolio 每只标的 1 行;multi 每个子策略 1 行 + (行内 ``strategy_name`` 是所属已保存策略的名字)。 + """ + + strategy_id: str + strategy_name: str + kind: Literal["single", "portfolio", "multi"] + strategy: str + strategy_label: str = "" + params: dict[str, Any] = {} + symbol: str + category: str = "DAY" + latest_signal: Literal["BUY", "SELL"] | None = None # 窗口内最后一根有信号的 K 线 + signal_date: str | None = None # 该信号所在 K 线日期 + recent_signals: list[SignalScanRecentSignal] = [] # 窗口内全部信号(按时间正序) + position: Literal["holding", "flat"] | None = None # 扫描结束时策略仓位 + last_close: float | None = None + last_bar_date: str | None = None + error: str | None = None + + +class SignalScanResult(BaseModel): + """信号扫描结果:全部子任务行 + 汇总计数。""" + + rows: list[SignalScanRow] + total: int = 0 + buy_count: int = 0 # 窗口内有买入信号的行数 + sell_count: int = 0 # 窗口内有卖出信号的行数 + error_count: int = 0 + elapsed: float = 0.0 + + # ── 结果序列化 ───────────────────────────────────────────────────────────────── diff --git a/src/easy_tdx/web/routers/backtest.py b/src/easy_tdx/web/routers/backtest.py index 345812b..8d6cfbd 100644 --- a/src/easy_tdx/web/routers/backtest.py +++ b/src/easy_tdx/web/routers/backtest.py @@ -25,6 +25,7 @@ from easy_tdx.web.backtest_schemas import ( OptimizeAllResult, OptimizeBacktestRequest, PortfolioBacktestRequest, + SignalScanRequest, StrategySchemaResponse, TaskListResponse, TaskStateResponse, @@ -297,6 +298,44 @@ async def run_optimize_all_async( return TaskSubmitResponse(task_id=task_id, status=status) +# ── 信号雷达(一键扫描已保存策略)──────────────────────────────────────────── + + +@router.post("/backtest/signal-scan/run/async", response_model=TaskSubmitResponse, status_code=202) +async def run_signal_scan_async( + req: SignalScanRequest, + client: Any = Depends(get_client), +) -> TaskSubmitResponse: + """提交「信号雷达」后台任务:扫描策略库全部已保存策略的最近买卖信号。 + + single/portfolio/multi 统一展开成"策略×标的"子任务,按 (symbol, category) + 去重取最近 800 根 K 线(async 上下文内完成),后台线程内逐条跑信号流程 + (与回测引擎同口径,含仓位跟踪)。只扫信号、不重跑回测、不改业绩快照。 + 结果为 SignalScanResult,通过 GET /backtest/tasks/{task_id} 轮询。 + """ + from easy_tdx.web.signal_scan import expand_targets, fetch_scan_bars, run_scan + from easy_tdx.web.strategy_store import get_store + + records = get_store().list_all() + if not records: + raise ValueError("策略库为空,请先在回测页保存策略") + + targets = expand_targets(records) + bars = await fetch_scan_bars(client, targets) + description = ( + f"信号扫描 | {len(records)}条策略 · {len(targets)}个子任务 · 窗口{req.window_bars}根" + ) + + runner = get_runner() + task_id = runner.submit( + lambda: run_scan(bars, targets, req.window_bars), + description=description, + ) + state = runner.get(task_id) + status: Any = state.status if state.status in ("pending", "running") else "running" + return TaskSubmitResponse(task_id=task_id, status=status) + + # ── 内部实现 ─────────────────────────────────────────────────────────────────── diff --git a/src/easy_tdx/web/signal_scan.py b/src/easy_tdx/web/signal_scan.py new file mode 100644 index 0000000..dbb4203 --- /dev/null +++ b/src/easy_tdx/web/signal_scan.py @@ -0,0 +1,325 @@ +"""信号雷达:一键扫描策略库全部已保存策略的最近买卖信号。 + +流程(与「策略库 → 重跑到今天」同一套信号口径): +1. ``expand_targets``: 把已保存策略(single/portfolio/multi 三种 kind)统一展开成 + "策略×标的" 子任务列表;数据损坏的条目展开为带 error 的行,不中断整批。 +2. ``fetch_scan_bars``: 按 (symbol, category) 去重取最近 K 线(async,event loop 内 + 调用;单页 800 根足够覆盖内置策略全部参数的指标预热)。 +3. ``run_scan``: 后台线程内逐 target 构建策略实例,跑一遍 bar-by-bar 信号流程 + (复用 combo._update_position 跟踪仓位,与 BacktestEngine 同口径), + 汇总最近 ``window`` 根内的买卖信号、结束仓位与最新收盘价。 + +只扫信号、不重跑完整回测,也不改写策略库保存的业绩快照。 +""" + +from __future__ import annotations + +import logging +import re +import time +from dataclasses import dataclass, field +from typing import Any + +import pandas as pd + +from easy_tdx.backtest.combo import _update_position +from easy_tdx.backtest.strategy import Strategy +from easy_tdx.web.strategy_store import SavedStrategy + +logger = logging.getLogger(__name__) + +# 每标的取的 K 线根数:标准协议单次上限 800 根,足够内置策略最慢参数(如慢线 250)预热。 +SCAN_BARS = 800 + +# 仓位跟踪用的佣金率(与 combo.extract_factor_signals 默认一致,只影响全仓股数估算) +_COMMISSION = 0.0003 + +# 市场前缀纠错规则(与前端 web-ui/src/market.ts detectMarket 保持一致): +# 北交所 43/83/87/92/93/4xx/8xx;沪市 6xx/9xx/5xx(含沪市基金);其余深市。 +_BJ_PREFIX = re.compile(r"^(43|83|87|92|93|4|8)") +_SH_PREFIX = re.compile(r"^[695]") + + +def _detect_market(code: str) -> str: + """按 6 位代码推断市场(SH/SZ/BJ),规则与前端 detectMarket 一致。""" + if not re.fullmatch(r"\d{6}", code): + return "SZ" + if _BJ_PREFIX.match(code): + return "BJ" + if _SH_PREFIX.match(code): + return "SH" + return "SZ" + + +def normalize_symbol(raw: str) -> str: + """纠正历史保存策略的市场前缀(如 SZ:515080 → SH:515080)。 + + 早期前端曾按市场前缀漏判沪市基金,导致部分历史保存的 symbol 错标, + 后端按错配市场取到 0 根 K 线被静默跳过。这里按代码段重判市场兜底。 + """ + code = raw.split(":", 1)[-1].strip() if raw else "" + if not code: + return raw + return f"{_detect_market(code)}:{code}" + + +# ── 展开子任务 ──────────────────────────────────────────────────────────────── + + +@dataclass +class ScanTarget: + """一个待扫描的"策略×标的"子任务(由已保存策略展开而来)。 + + ``error`` 非空表示展开阶段就发现问题(缺 symbol / 组合数据损坏), + run_scan 会把它原样写进结果行,不参与取数与信号计算。 + """ + + strategy_id: str # 所属已保存策略 id + strategy_name: str # 所属已保存策略名(展示用) + kind: str # single | portfolio | multi + strategy: str # 策略注册表 key(如 ma_cross) + strategy_label: str = "" + params: dict[str, Any] = field(default_factory=dict) + symbol: str = "" # 归一化后的 "市场:代码" + category: str = "DAY" + error: str | None = None + + +def expand_targets(records: list[SavedStrategy]) -> list[ScanTarget]: + """把全部已保存策略展开成"策略×标的"子任务列表。 + + - single: 1 条(context.symbol) + - portfolio: context.stocks 每只一条(同 strategy + params) + - multi: context.items 每条一 target(各自带 strategy/params/symbol) + - 缺关键字段的条目展开为 error 行(保证结果表能看到"这条策略有问题") + """ + targets: list[ScanTarget] = [] + for rec in records: + ctx = rec.context or {} + if rec.kind == "multi": + items = ctx.get("items") + if not isinstance(items, list) or not items: + targets.append( + ScanTarget( + strategy_id=rec.id, + strategy_name=rec.name, + kind=rec.kind, + strategy=rec.strategy, + error="组合缺少策略明细(items),可能数据损坏", + ) + ) + continue + for item in items: + if not isinstance(item, dict): + targets.append(_error_target(rec, "组合条目数据损坏")) + continue + symbol = item.get("symbol") + if not item.get("strategy") or not symbol: + targets.append(_error_target(rec, "组合条目缺少 strategy/symbol")) + continue + targets.append( + ScanTarget( + strategy_id=rec.id, + strategy_name=rec.name, + kind=rec.kind, + strategy=str(item["strategy"]), + strategy_label=str(item.get("strategy_label") or ""), + params=item.get("params") or {}, + symbol=normalize_symbol(str(symbol)), + category=str(item.get("category") or "DAY"), + ) + ) + else: + # single 与 portfolio 同构:portfolio 把同策略铺到多只标的 + stocks = ctx.get("stocks") if rec.kind == "portfolio" else None + symbols = [str(s) for s in stocks] if isinstance(stocks, list) and stocks else None + if symbols is None: + symbol = ctx.get("symbol") + if not symbol: + targets.append(_error_target(rec, "缺少标的上下文(symbol)")) + continue + symbols = [str(symbol)] + for sym in symbols: + targets.append( + ScanTarget( + strategy_id=rec.id, + strategy_name=rec.name, + kind=rec.kind, + strategy=rec.strategy, + strategy_label=rec.strategy_label, + params=rec.params or {}, + symbol=normalize_symbol(sym), + category=str(ctx.get("category") or "DAY"), + ) + ) + return targets + + +def _error_target(rec: SavedStrategy, message: str) -> ScanTarget: + """构造一条展开失败的 error 行(保留策略身份,便于在结果表定位)。""" + return ScanTarget( + strategy_id=rec.id, + strategy_name=rec.name, + kind=rec.kind, + strategy=rec.strategy, + error=message, + ) + + +# ── 取行情 ──────────────────────────────────────────────────────────────────── + + +async def fetch_scan_bars( + client: Any, + targets: list[ScanTarget], +) -> dict[tuple[str, str], pd.DataFrame | None]: + """按 (symbol, category) 去重取最近 ``SCAN_BARS`` 根 K 线(async,event loop 内调用)。 + + 同一标的被多个策略引用时只取一次。单个标的取数失败/数据无效记 None + (不中断整批),run_scan 会给相关行统一标 error。 + """ + from easy_tdx.web.convert import category_from_str, market_from_str + + bars: dict[tuple[str, str], pd.DataFrame | None] = {} + for t in targets: + key = (t.symbol, t.category) + if key in bars or t.error: + continue + try: + market_str, code = t.symbol.split(":", 1) + df = await client.get_security_bars( + market_from_str(market_str), + code, + category_from_str(t.category), + 0, + SCAN_BARS, + ) + except Exception as exc: # noqa: BLE001 — 单标的失败不中断整批 + logger.warning("信号扫描取数失败 %s: %s", t.symbol, exc) + bars[key] = None + continue + if not isinstance(df, pd.DataFrame) or len(df) < 2 or "close" not in df.columns: + bars[key] = None + continue + # 列归一化:日线返回 date 列,_bind_data 需要 datetime;页内已正序但保险再排一次 + if "datetime" not in df.columns and "date" in df.columns: + df = df.copy() + df["datetime"] = df["date"] + if "datetime" not in df.columns: + bars[key] = None + continue + bars[key] = df.sort_values("datetime").reset_index(drop=True) + return bars + + +# ── 信号评估 ────────────────────────────────────────────────────────────────── + + +def evaluate_signals( + strategy: Strategy, + df: pd.DataFrame, + window: int, +) -> dict[str, Any]: + """在 df 上单遍跑策略的 bar-by-bar 信号流程,返回最近 window 根内的信号摘要。 + + 复现 BacktestEngine._generate_signals / combo.extract_factor_signals 的 + 信号收集 + 仓位跟踪(``_update_position``),保证扫描结果与真实回测一致。 + """ + n = len(df) + strat = strategy + strat._bind_data(df) + strat._cash = 100_000.0 + strat._position_size = 0.0 + strat._call_init() + + close_arr = df["close"].to_numpy() + dt_col = df["datetime"] + start = max(0, n - window) + recent: list[dict[str, Any]] = [] + + for i in range(n): + strat._set_bar_index(i) + strat._call_next() + signals = strat._clear_signals() + if signals and i >= start: + date = str(dt_col.iloc[i])[:16] + for sig in signals: + recent.append({"date": date, "direction": sig.direction}) + _update_position(strat, signals, close_arr[i], _COMMISSION) + + return { + "recent_signals": recent, + "latest_signal": recent[-1]["direction"] if recent else None, + "signal_date": recent[-1]["date"] if recent else None, + # 结束仓位(容忍浮点误差):>0 视为策略当前持仓 + "position": "holding" if strat._position_size > 0.5 else "flat", + "last_close": float(close_arr[-1]), + "last_bar_date": str(dt_col.iloc[-1])[:16], + } + + +# ── 汇总扫描 ────────────────────────────────────────────────────────────────── + + +def run_scan( + bars: dict[tuple[str, str], pd.DataFrame | None], + targets: list[ScanTarget], + window: int, +) -> dict[str, Any]: + """后台线程内执行:逐 target 构建策略实例并评估信号,汇总成扫描结果。 + + 单个 target 失败(未知策略/参数非法/取数为空/计算异常)记为该行的 + error,不影响其余行。返回结构对应 SignalScanResult schema。 + """ + from easy_tdx.backtest.strategies import get_registry + + registry = get_registry() + rows: list[dict[str, Any]] = [] + t0 = time.time() + + for t in targets: + row: dict[str, Any] = { + "strategy_id": t.strategy_id, + "strategy_name": t.strategy_name, + "kind": t.kind, + "strategy": t.strategy, + "strategy_label": t.strategy_label, + "params": t.params, + "symbol": t.symbol, + "category": t.category, + "latest_signal": None, + "signal_date": None, + "recent_signals": [], + "position": None, + "last_close": None, + "last_bar_date": None, + "error": None, + } + try: + if t.error: + raise ValueError(t.error) + df = bars.get((t.symbol, t.category)) + if df is None: + raise ValueError("未取到有效 K 线(停牌/代码失效/取数失败)") + try: + entry = registry.get(t.strategy) + except KeyError as exc: + raise ValueError(f"未知策略 '{t.strategy}'(可能为旧版本保存)") from exc + strategy = entry.build(t.params) + row.update(evaluate_signals(strategy, df, window)) + except Exception as exc: # noqa: BLE001 — 单行失败不中断整批 + row["error"] = str(exc) or type(exc).__name__ + logger.warning("信号扫描失败 %s@%s: %s", t.strategy, t.symbol, exc) + rows.append(row) + + buy_count = sum(1 for r in rows if any(s["direction"] == "BUY" for s in r["recent_signals"])) + sell_count = sum(1 for r in rows if any(s["direction"] == "SELL" for s in r["recent_signals"])) + error_count = sum(1 for r in rows if r["error"]) + return { + "rows": rows, + "total": len(rows), + "buy_count": buy_count, + "sell_count": sell_count, + "error_count": error_count, + "elapsed": round(time.time() - t0, 2), + } diff --git a/tests/unit/test_signal_scan.py b/tests/unit/test_signal_scan.py new file mode 100644 index 0000000..918b544 --- /dev/null +++ b/tests/unit/test_signal_scan.py @@ -0,0 +1,386 @@ +"""信号雷达(signal_scan)单元 + 端到端测试(离线,无网络)。 + +覆盖: +- normalize_symbol 市场前缀纠错 +- expand_targets 三种 kind 展开 + 数据损坏容错 +- fetch_scan_bars 去重取数 / 失败容错 / date→datetime 列归一化 +- evaluate_signals 金叉买入、死叉卖出、仓位跟踪、窗口过滤(与回测引擎同口径) +- run_scan 单行失败不中断 + 汇总计数 +- POST /backtest/signal-scan/run/async 端到端(fake store + fake 行情) +""" + +from __future__ import annotations + +import asyncio +import time + +import numpy as np +import pandas as pd +import pytest + +pytest.importorskip("fastapi") + +from easy_tdx.web.signal_scan import ( # noqa: E402 + evaluate_signals, + expand_targets, + fetch_scan_bars, + normalize_symbol, + run_scan, +) +from easy_tdx.web.strategy_store import SavedStrategy # noqa: E402 + +# ── 测试数据 ─────────────────────────────────────────────────────────────────── + + +def v_shape_df(n_fall: int = 40, n_rise: int = 80, n_drop: int = 0) -> pd.DataFrame: + """V 型走势合成日线:下跌 → 上涨(→ 可选急跌),保证出现金叉(→ 死叉)。 + + 返回的 df 带标准 OHLCV + datetime 列(日线接口返回 date,归一化后是 datetime)。 + """ + closes = np.concatenate( + [ + 10.0 - np.arange(n_fall) * 0.02, # 缓跌:MA5 持续低于 MA20 + 9.2 + np.arange(n_rise) * 0.12, # 稳定上涨:金叉出现 + (10.0 + n_rise * 0.12 - np.arange(1, n_drop + 1) * 0.5) if n_drop else [], # 急跌:死叉 + ] + ) + n = len(closes) + dates = pd.date_range("2025-01-01", periods=n, freq="B") + return pd.DataFrame( + { + "datetime": dates, + "open": closes - 0.05, + "high": closes + 0.10, + "low": closes - 0.10, + "close": closes, + "vol": np.full(n, 5000.0), + "amount": closes * 5000, + } + ) + + +def _single(**ctx_overrides: object) -> SavedStrategy: + ctx: dict = {"symbol": "SH:601088", "category": "DAY"} + ctx.update(ctx_overrides) + return SavedStrategy( + id="s1", + name="神华·双均线", + kind="single", + strategy="ma_cross", + strategy_label="双均线交叉", + params={"fast": 5, "slow": 20}, + context=ctx, + ) + + +class FakeClient: + """假行情客户端:按 (market, code) 返回预置 df,未预置的抛错。""" + + def __init__(self, data: dict[str, pd.DataFrame]) -> None: + self.data = data + self.calls: list[tuple[str, str]] = [] + + async def get_security_bars(self, market, code, category, start, count): # noqa: ANN001 + market_str = str(getattr(market, "name", market)) + key = f"{market_str}:{code}" + self.calls.append((key, str(getattr(category, "name", category)))) + if key not in self.data: + raise ConnectionError(f"no data for {key}") + return self.data[key] + + +# ── normalize_symbol ────────────────────────────────────────────────────────── + + +@pytest.mark.parametrize( + ("raw", "expected"), + [ + ("SH:601088", "SH:601088"), # 正确的沪市主板 + ("SZ:515080", "SH:515080"), # 历史错标的沪市基金 → 纠正 + ("510300", "SH:510300"), # 无前缀 → 补全 + ("SZ:000001", "SZ:000001"), # 正确的深市主板 + ("430047", "BJ:430047"), # 北交所 + ("830799", "BJ:830799"), # 北交所 8xx + ("SZ:300347", "SZ:300347"), # 创业板 + ], +) +def test_normalize_symbol(raw: str, expected: str) -> None: + assert normalize_symbol(raw) == expected + + +# ── expand_targets ──────────────────────────────────────────────────────────── + + +def test_expand_single() -> None: + targets = expand_targets([_single()]) + assert len(targets) == 1 + t = targets[0] + assert (t.strategy, t.params, t.symbol, t.category) == ( + "ma_cross", + {"fast": 5, "slow": 20}, + "SH:601088", + "DAY", + ) + assert t.error is None + + +def test_expand_portfolio_multi_symbols() -> None: + rec = SavedStrategy( + id="p1", + name="银行组合", + kind="portfolio", + strategy="macd", + params={"short": 10, "long": 20}, + context={"stocks": ["SZ:000001", "515080", "SH:601088"], "category": "DAY"}, + ) + targets = expand_targets([rec]) + assert len(targets) == 3 + assert [t.symbol for t in targets] == ["SZ:000001", "SH:515080", "SH:601088"] + assert all(t.strategy == "macd" for t in targets) + + +def test_expand_multi_items() -> None: + rec = SavedStrategy( + id="m1", + name="老登+小登组合", + kind="multi", + strategy="multi", + context={ + "items": [ + {"strategy": "trix", "params": {"m1": 18, "m2": 20}, "symbol": "SZ:300347"}, + { + "strategy": "ema_cross", + "params": {"fast": 12}, + "symbol": "SZ:301308", + "category": "DAY", + }, + {"strategy": "macd", "symbol": "SH:601088"}, # 缺 params → 默认空 + {"strategy": "", "symbol": "SZ:000001"}, # 缺 strategy → error 行 + ] + }, + ) + targets = expand_targets([rec]) + assert len(targets) == 4 + ok = [t for t in targets if t.error is None] + assert [t.strategy for t in ok] == ["trix", "ema_cross", "macd"] + assert ok[1].params == {"fast": 12} + assert [t.error is None for t in targets] == [True, True, True, False] + + +def test_expand_error_rows() -> None: + # single 缺 symbol / multi 缺 items → 各展开为一条 error 行(不丢策略身份) + no_symbol = _single() + no_symbol.context = {"category": "DAY"} + broken_multi = SavedStrategy(id="m2", name="坏组合", kind="multi", strategy="multi") + targets = expand_targets([no_symbol, broken_multi]) + assert len(targets) == 2 + assert all(t.error for t in targets) + assert [t.strategy_name for t in targets] == ["神华·双均线", "坏组合"] + + +# ── fetch_scan_bars ─────────────────────────────────────────────────────────── + + +def test_fetch_scan_bars_dedupe_and_normalize() -> None: + df = v_shape_df() + # 日线接口风格:date 列而非 datetime + daily = df.rename(columns={"datetime": "date"}) + client = FakeClient({"SH:601088": daily, "SZ:000001": daily}) + rec1 = _single() + rec2 = _single(id="s2", name="另一个神华", strategy="macd", params={}) + targets = expand_targets([rec1, rec2]) # 同 symbol 只取一次 + targets.append(expand_targets([_single(symbol="SZ:000001")])[0]) + + bars = asyncio.run(fetch_scan_bars(client, targets)) + assert set(bars) == {("SH:601088", "DAY"), ("SZ:000001", "DAY")} + # SH:601088 只取了一次(去重生效) + assert len([c for c in client.calls if c[0] == "SH:601088"]) == 1 + # date 列已归一化为 datetime 且按时间正序 + out = bars[("SH:601088", "DAY")] + assert "datetime" in out.columns + assert out["datetime"].is_monotonic_increasing + + +def test_fetch_scan_bars_failure_tolerant() -> None: + client = FakeClient({}) # 全部抛错 + targets = expand_targets([_single()]) + bars = asyncio.run(fetch_scan_bars(client, targets)) + assert bars == {("SH:601088", "DAY"): None} + + +# ── evaluate_signals ────────────────────────────────────────────────────────── + + +def _ma_cross_instance(): # noqa: ANN202 + from easy_tdx.backtest.strategies import get_registry + + return get_registry().get("ma_cross").build({"fast": 5, "slow": 20}) + + +def _expected_cross_dates(df: pd.DataFrame, direction: str) -> list[str]: + """用 MyTT 独立算出金叉/死叉所在日期(作为期望值,与被测代码解耦)。""" + from easy_tdx.MyTT import CROSS, MA + + close = df["close"].to_numpy() + fast, slow = MA(close, 5), MA(close, 20) + mask = CROSS(fast, slow) if direction == "BUY" else CROSS(slow, fast) + return [str(df["datetime"].iloc[i])[:16] for i in range(len(df)) if mask[i]] + + +def test_evaluate_signals_golden_cross_buy() -> None: + df = v_shape_df(n_rise=30) # 只涨不跌:恰好一个金叉、之后无死叉 + buy_dates = _expected_cross_dates(df, "BUY") + assert len(buy_dates) == 1, "V 型数据应恰好产生一个金叉" + cross_date = buy_dates[0] + cross_idx = [i for i in range(len(df)) if str(df["datetime"].iloc[i])[:16] == cross_date][0] + + # 窗口恰好从金叉那根开始 → 窗口内能捕获 BUY + result = evaluate_signals(_ma_cross_instance(), df, window=len(df) - cross_idx) + buys = [s for s in result["recent_signals"] if s["direction"] == "BUY"] + assert [s["date"] for s in buys] == [cross_date] + assert result["latest_signal"] == "BUY" + assert result["signal_date"] == cross_date + assert result["position"] == "holding" # 买入后一直持有 + assert result["last_close"] == pytest.approx(float(df["close"].iloc[-1])) + assert result["last_bar_date"] == str(df["datetime"].iloc[-1])[:16] + + # 窗口再收窄一根(金叉在窗口外)→ 不上报旧信号,但仓位跟踪不受窗口影响 + result2 = evaluate_signals(_ma_cross_instance(), df, window=len(df) - cross_idx - 1) + assert result2["recent_signals"] == [] + assert result2["latest_signal"] is None + assert result2["position"] == "holding" + + +def test_evaluate_signals_death_cross_sell() -> None: + df = v_shape_df(n_drop=15) # 涨完急跌:金叉买入 → 死叉卖出 + result = evaluate_signals(_ma_cross_instance(), df, window=len(df)) + dirs = [s["direction"] for s in result["recent_signals"]] + assert dirs[0] == "BUY" + assert dirs[-1] == "SELL" + assert result["latest_signal"] == "SELL" + assert result["position"] == "flat" # 清仓 + + +def test_evaluate_signals_matches_engine_trades() -> None: + """与真实回测引擎成交方向序列一致性抽查(同 df、同策略)。""" + from easy_tdx.backtest.engine import BacktestEngine + + df = v_shape_df(n_drop=15) + strat = _ma_cross_instance() + result = evaluate_signals(strat, df, window=len(df)) + engine = BacktestEngine(strategy=_ma_cross_instance()) + trades = engine.run(df).trades + engine_dirs = list(trades["direction"]) + scan_dirs = [s["direction"] for s in result["recent_signals"]] + assert scan_dirs == engine_dirs[: len(scan_dirs)] + + +# ── run_scan ────────────────────────────────────────────────────────────────── + + +def test_run_scan_summary_and_errors() -> None: + df = v_shape_df() + targets = [ + expand_targets([_single()])[0], # 正常行(有行情) + expand_targets([_single(id="s2", name="同标的第二策略")])[0], # 同标的复用行情 + ] + # 制造三类失败:未知策略 / 无行情 / 展开错误 + bad_strategy = expand_targets([_single()])[0] + bad_strategy.strategy = "nope_strategy" + targets.append(bad_strategy) + no_bars = expand_targets([_single()])[0] + no_bars.symbol = "SZ:999999" + targets.append(no_bars) + broken = expand_targets([SavedStrategy(id="x", name="坏", kind="single", strategy="ma_cross")])[ + 0 + ] + targets.append(broken) + + bars = {("SH:601088", "DAY"): df} + out = run_scan(bars, targets, window=len(df)) + assert out["total"] == 5 + assert out["buy_count"] == 2 # 前两行各有一个金叉买入 + assert out["sell_count"] == 0 + assert out["error_count"] == 3 # 未知策略 / 无行情 / 展开错误 + rows = out["rows"] + assert rows[0]["error"] is None + assert rows[0]["latest_signal"] == "BUY" + assert "未知策略" in rows[2]["error"] + assert "未取到有效 K 线" in rows[3]["error"] + assert "缺少标的上下文" in rows[4]["error"] + assert out["elapsed"] >= 0 + + +# ── API 端到端 ──────────────────────────────────────────────────────────────── + + +@pytest.fixture() +def api_client(): + from fastapi.testclient import TestClient + + from easy_tdx.web import create_app + + app = create_app() + with TestClient(app) as c: + yield c + + +def test_signal_scan_endpoint_e2e(api_client, monkeypatch) -> None: + """POST 提交 → 轮询 done → 结果结构完整(fake store + fake 取数)。""" + import easy_tdx.web.signal_scan as sigscan + import easy_tdx.web.strategy_store as store_mod + + # 缓跌 59 根 + 末根跳涨:金叉恰好发生在最后一根 K 线(窗口=1 也能捕获) + df = v_shape_df(n_fall=59, n_rise=0) + df.loc[df.index[-1], ["open", "high", "low", "close"]] = [14.95, 15.2, 14.8, 15.0] + + class FakeStore: + def list_all(self) -> list[SavedStrategy]: + return [_single()] + + async def fake_fetch(client, targets): # noqa: ANN001 + return {("SH:601088", "DAY"): df} + + monkeypatch.setattr(store_mod, "get_store", lambda: FakeStore()) + monkeypatch.setattr(sigscan, "fetch_scan_bars", fake_fetch) + + resp = api_client.post("/api/v1/backtest/signal-scan/run/async", json={"window_bars": 1}) + assert resp.status_code == 202, resp.text + task_id = resp.json()["task_id"] + + final = None + for _ in range(200): + poll = api_client.get(f"/api/v1/backtest/tasks/{task_id}") + assert poll.status_code == 200 + final = poll.json() + if final["status"] in ("done", "failed"): + break + time.sleep(0.05) + assert final is not None and final["status"] == "done", final + + result = final["result"] + assert result["total"] == 1 + assert result["buy_count"] == 1 + row = result["rows"][0] + assert row["strategy"] == "ma_cross" + assert row["symbol"] == "SH:601088" + assert row["error"] is None + assert row["position"] in ("holding", "flat") + + +def test_signal_scan_endpoint_empty_store(api_client, monkeypatch) -> None: + import easy_tdx.web.strategy_store as store_mod + + class EmptyStore: + def list_all(self) -> list[SavedStrategy]: + return [] + + monkeypatch.setattr(store_mod, "get_store", lambda: EmptyStore()) + resp = api_client.post("/api/v1/backtest/signal-scan/run/async", json={}) + assert resp.status_code == 400 + assert "策略库为空" in resp.json()["detail"] + + +def test_signal_scan_endpoint_window_validation(api_client) -> None: + resp = api_client.post("/api/v1/backtest/signal-scan/run/async", json={"window_bars": 0}) + assert resp.status_code == 422 diff --git a/web-ui/src/App.vue b/web-ui/src/App.vue index 8ab88a2..64cd7de 100644 --- a/web-ui/src/App.vue +++ b/web-ui/src/App.vue @@ -12,6 +12,7 @@ 参数寻优 结果对比 策略库 + 信号雷达 服务器设置 diff --git a/web-ui/src/api.ts b/web-ui/src/api.ts index 08ba6ea..d1fecd2 100644 --- a/web-ui/src/api.ts +++ b/web-ui/src/api.ts @@ -17,6 +17,8 @@ import type { ServerHostInfo, ServerHostListResponse, ServerSwitchResult, + SignalScanRequest, + SignalScanResult, StrategiesResponse, TaskListResponse, TaskState, @@ -237,6 +239,55 @@ export async function runBacktestWithPolling( // ── 策略库(已保存策略)────────────────────────────────────────────────────── +/** 提交「信号雷达」一键扫描后台任务,返回 task_id。 */ +export async function submitSignalScanTask( + req: SignalScanRequest = {}, +): Promise { + const resp = await fetch(`${BASE}/backtest/signal-scan/run/async`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(req), + }) + if (!resp.ok) await throwError(resp) + return (await resp.json()) as TaskSubmitResponse +} + +/** + * 提交信号扫描并轮询直到 done/failed。 + * + * 与 runBacktestWithPolling 的区别:扫描要在请求内逐标的取行情(提交本身 + * 就可能耗时数十秒),且标的较多时总时长可能超过 2 分钟,故默认 300s 超时。 + */ +export async function runSignalScanWithPolling( + req: SignalScanRequest = {}, + onPoll?: (state: TaskState) => void, + intervalMs = 500, + timeoutMs = 300_000, +): Promise { + const { task_id } = await submitSignalScanTask(req) + const start = Date.now() + // eslint-disable-next-line no-constant-condition + while (true) { + const state = await fetchTask(task_id) + onPoll?.(state) + if (state.status === 'done' || state.status === 'failed') return state + if (Date.now() - start > timeoutMs) { + throw new Error(`信号扫描超时(${timeoutMs / 1000}s),可稍后重试或减小窗口`) + } + await new Promise((r) => setTimeout(r, intervalMs)) + } +} + +/** 断言任务结果为信号扫描结果(类型收窄用)。 */ +export function asSignalScanResult(state: TaskState): SignalScanResult { + if (state.status === 'failed') throw new Error(state.error || '信号扫描失败') + const result = state.result as SignalScanResult | null + if (!result || !Array.isArray(result.rows)) { + throw new Error('信号扫描结果格式异常(缺少 rows)') + } + return result +} + /** 列出全部已保存策略(按创建时间倒序)。 */ export async function fetchSavedStrategies(): Promise { const resp = await fetch(`${BASE}/strategies`) diff --git a/web-ui/src/router.ts b/web-ui/src/router.ts index c2a1575..4a12be5 100644 --- a/web-ui/src/router.ts +++ b/web-ui/src/router.ts @@ -5,15 +5,18 @@ import CompareView from './views/CompareView.vue' import OptimizeView from './views/OptimizeView.vue' import PortfolioView from './views/PortfolioView.vue' import ServerSettingsView from './views/ServerSettingsView.vue' +import SignalRadarView from './views/SignalRadarView.vue' import StrategiesView from './views/StrategiesView.vue' -// 单标的回测(/)+ 组合回测(/portfolio)+ 参数寻优(/optimize)+ 结果对比(/compare)+ 策略库(/strategies)+ 服务器设置(/settings)。 +// 单标的回测(/)+ 组合回测(/portfolio)+ 参数寻优(/optimize)+ 结果对比(/compare) +// + 策略库(/strategies)+ 信号雷达(/signals)+ 服务器设置(/settings)。 const routes = [ { path: '/', name: 'backtest', component: BacktestView }, { path: '/portfolio', name: 'portfolio', component: PortfolioView }, { path: '/optimize', name: 'optimize', component: OptimizeView }, { path: '/compare', name: 'compare', component: CompareView }, { path: '/strategies', name: 'strategies', component: StrategiesView }, + { path: '/signals', name: 'signals', component: SignalRadarView }, { path: '/settings', name: 'settings', component: ServerSettingsView }, ] diff --git a/web-ui/src/types.ts b/web-ui/src/types.ts index f1a5ea2..4d5b7e8 100644 --- a/web-ui/src/types.ts +++ b/web-ui/src/types.ts @@ -131,7 +131,13 @@ export type TaskStatus = 'pending' | 'running' | 'done' | 'failed' export interface TaskState { task_id: string status: TaskStatus - result: BacktestResult | PortfolioResult | OptimizeResult | OptimizeAllResult | null + result: + | BacktestResult + | PortfolioResult + | OptimizeResult + | OptimizeAllResult + | SignalScanResult + | null error: string | null description: string elapsed: number @@ -307,6 +313,48 @@ export interface SavedStrategyListResponse { count: number } +// ── 信号雷达(POST /api/v1/backtest/signal-scan/run/async) ────────────────── + +/** 信号扫描请求:window_bars = 检查最近 N 根 K 线内的信号。 */ +export interface SignalScanRequest { + window_bars?: number +} + +/** 窗口内单根 K 线的信号。 */ +export interface SignalScanRecentSignal { + date: string + direction: 'BUY' | 'SELL' +} + +/** 扫描结果单行:一个"策略×标的"子任务的信号摘要。 */ +export interface SignalScanRow { + strategy_id: string + strategy_name: string + kind: 'single' | 'portfolio' | 'multi' + strategy: string + strategy_label: string + params: Record + symbol: string + category: string + latest_signal: 'BUY' | 'SELL' | null + signal_date: string | null + recent_signals: SignalScanRecentSignal[] + position: 'holding' | 'flat' | null + last_close: number | null + last_bar_date: string | null + error: string | null +} + +/** 信号扫描结果:全部子任务行 + 汇总计数。 */ +export interface SignalScanResult { + rows: SignalScanRow[] + total: number + buy_count: number + sell_count: number + error_count: number + elapsed: number +} + // ── 多策略组合回测(资金分仓,POST /api/v1/backtest/multi-strategy/run/async) ── /** 多策略组合的单个策略槽位(一个策略 + 参数 + 它要跑的原标的 + 日期)。 */ diff --git a/web-ui/src/views/SignalRadarView.vue b/web-ui/src/views/SignalRadarView.vue new file mode 100644 index 0000000..1f94b30 --- /dev/null +++ b/web-ui/src/views/SignalRadarView.vue @@ -0,0 +1,587 @@ + + + + +