mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 23:44:16 +08:00
* fix(concurrency): 共享缓存/任务表加锁, 全局限速, depth 原子写, 认证热路径缓存 修复多线程下的竞态与阻塞: - overview/strategy_cache/PanelCache/StrategyMonitor._watching 四处共享状态加锁, 消除 "dict/OrderedDict mutated" 与丢更新/半写读取 - strategy_cache/depth parquet 改临时文件 + os.replace 原子写 - rate_limits 改进程级共享时间轴限速, 并发同步不再聚合超过单能力 rpm; scheduler 令牌账目与 sleep 分离, sleep 不再独占锁串行化其他请求 - auth.is_configured() 内存缓存, 认证中间件不再每请求读盘阻塞事件循环 - api/backtest 任务清理/取消全程持 _jobs_lock, 并用 Semaphore(2) 限并发重回测 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * perf(data): limit_ladder 去 N+1 全市场重算, 指标裁剪, factor 向量化 - limit_ladder 前一日 consecutive 改窄读单日 parquet 存储列 (谓词/投影下推), 替代 range(1,10) 逐日 _load_enriched_for_date 全市场指标重算 (最坏 9x) - compute_indicators 新增可选 needed 裁剪 (默认 None 行为逐位不变, 已对照验证), factor 只算所需因子列 - factor._calc_period_return 用 Polars join 替代 Python 逐行 price_map 循环, _add_groups 去 map_elements 改纯表达式 (输出逐位一致) - screener ext value_map 按 parquet mtime 记忆化, 免每请求磁盘重读 (DuckDB 过滤仍用隔离 :memory: 连接, 不扩大注入面) Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * refactor(backend): 报表存储去重, 删死代码, DuckDB 视图重建收敛, 管道失败如实标记 - 三份近乎逐字复制的 *_reports.py 收敛到共享 JsonReportStore (原子写 + 锁), 各模块公有 API/id 格式/上限/落盘 schema 完全保持不变 - 删除 ext_pull.py 中字节相同的死 _run_loop (Python 只绑第二个) 及无用 import - 13 张 DuckDB 视图重建收敛为唯一权威 repository.rebuild_views(), daily_pipeline 与 /api/data/clear 改为调用 (修好 clear 路径漏挂视图的漂移) - daily_pipeline 累积 stage_errors 并在末尾抛出, 部分失败不再误报成功; free/None 模式的能力门控跳过不计入失败 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * feat(frontend): SSE 连接态, 路由代码分割, 查询失效修复, 三态与无障碍 - 实时行情 SSE: 连接态 store + 指数退避 + 断线徽标/toast (避免静默丢告警); 回测 SSE 断线有界重连 + 可重试, 不再永久卡住进度条 - router 全部 React.lazy + Suspense, vite manualChunks 拆图表库 (echarts 变独立 1MB 懒加载 chunk, 首屏包显著减小) - 修 Data 清库后其它页显示旧数据 (改回广域失效); 修 Watchlist kline 失效键 永不匹配; query key 收敛到 QK 工厂 (新增 strategyDetail) - Monitor/Analysis/StockAnalysis/ExtPages/CustomSignals 补 loading 门控与 error/empty 三态区分 - 新增共享 Modal 原语 (焦点陷阱/ESC/焦点还原/aria), 改造 3 个高频弹窗; Toast/AlertToast 加 aria-live 与键盘可达; Watchlist/LimitUpLadder 卡片 memo 修复本轮 review 发现的缺陷: - Modal 焦点 effect 依赖 onClose 致每次输入抢焦点 → 改 ref 只装一次 - StrategySettingsDialog 删除确认框被 Modal 面板裁剪 → 移出作兄弟节点 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * fix(quant): 修正 ST 板块限价套错 与 因子 Sharpe 年化频率 两个不报错但会算错数的领域 bug: 1. ST 5% 涨跌停限幅被无条件套到创业板/科创板 ST 股: 注册制改革后 创业板(300/301)、科创板(688/689) 的风险警示股仍执行 20%, 北交所 30%, 只有主板 ST 才是 5%。原代码 _is_st 先判且覆盖板块限幅, 导致 创业板/科创板 ST 的涨停价按 5% 计算 → +5% 被误报涨停、真 +20% 涨停被漏报, 污染 signal_limit_up / consecutive_limit_ups / 连板梯队 / near_limit_up。 修正: ST 5% 仅在 ~(创业板|科创板|北交所) 时生效 (EOD + 盘中两条路径 + near_limit_up)。 2. 因子回测 Sharpe 一律乘 √252, 但 group_nav 每点是一个调仓周期收益: 月频调仓下是月收益, 乘 √252 会把 Sharpe 高估 √(252/12) ≈ 4.6x (周频 ≈2.2x), 使无效因子显示成明星因子, 废掉"先筛无效指标"的用途。 修正: 年化系数按 config.rebalance 取 √252/√52/√12。 新增 tests/test_st_limit_and_sharpe.py (5 例) 覆盖两处修正。 Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
514 lines
18 KiB
Python
514 lines
18 KiB
Python
"""回测 API — 信号回测 + 因子回测 + 策略回测。"""
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
import queue
|
|
import threading
|
|
from dataclasses import asdict
|
|
from datetime import date, timedelta
|
|
from typing import Literal
|
|
|
|
from fastapi import APIRouter, HTTPException, Request
|
|
from fastapi.responses import StreamingResponse
|
|
from pydantic import BaseModel, Field
|
|
|
|
from app.config import settings
|
|
from app.services.backtest import (
|
|
BacktestConfig,
|
|
BacktestService,
|
|
VectorbtUnavailable,
|
|
is_available,
|
|
)
|
|
|
|
router = APIRouter(prefix="/api/backtest", tags=["backtest"])
|
|
|
|
FACTOR_DEFAULT_DAYS = 180
|
|
STRATEGY_DEFAULT_DAYS = 365 * 3
|
|
BACKTEST_MAX_SERVER_DAYS = 186
|
|
FACTOR_MAX_SYMBOLS = 1000
|
|
BACKTEST_SERVER_GUARD_MESSAGE = (
|
|
"当前服务器内存约 1.8GB,回测区间最多支持 6 个月;"
|
|
"更长周期容易触发 OOM,建议在 8GB 以上内存环境或本机运行。"
|
|
)
|
|
|
|
|
|
def _get_engine(request: Request):
|
|
"""获取或创建 BacktestEngine (单例,PanelCache 跨请求生效)。"""
|
|
from app.backtest.engine import BacktestEngine
|
|
engine = getattr(request.app.state, "backtest_engine", None)
|
|
if engine is None:
|
|
engine = BacktestEngine(request.app.state.repo)
|
|
request.app.state.backtest_engine = engine
|
|
return engine
|
|
|
|
|
|
def _resolve_start(req: BaseModel, end: date, default_days: int) -> date:
|
|
"""未传 start 使用默认区间;显式传 null/空值表示全部历史。"""
|
|
start = getattr(req, "start")
|
|
if start is not None:
|
|
return start
|
|
if "start" in req.model_fields_set:
|
|
return date(1900, 1, 1)
|
|
return end - timedelta(days=default_days)
|
|
|
|
|
|
def _guard_server_backtest_range(start: date, end: date):
|
|
if not settings.backtest_range_guard:
|
|
return
|
|
days = (end - start).days + 1
|
|
if days > BACKTEST_MAX_SERVER_DAYS:
|
|
raise HTTPException(status_code=400, detail=BACKTEST_SERVER_GUARD_MESSAGE)
|
|
|
|
|
|
# ================================================================
|
|
# 状态
|
|
# ================================================================
|
|
|
|
@router.get("/status")
|
|
def status():
|
|
"""前端可用此接口判断回测页是否要灰显。"""
|
|
return {"available": True}
|
|
|
|
|
|
# ================================================================
|
|
# 信号回测 (现有接口,保持不变)
|
|
# ================================================================
|
|
|
|
class BacktestRequest(BaseModel):
|
|
symbols: list[str] = Field(..., min_length=1)
|
|
start: date | None = None
|
|
end: date | None = None
|
|
entries: list[str] = []
|
|
exits: list[str] = []
|
|
stop_loss_pct: float | None = None
|
|
max_hold_days: int | None = None
|
|
fees_pct: float = 0.0002
|
|
slippage_bps: float = 5
|
|
matching: Literal["close_t", "open_t+1"] = "close_t"
|
|
asset_type: str = "stock"
|
|
|
|
|
|
@router.post("/run")
|
|
def run(req: BacktestRequest, request: Request):
|
|
"""信号回测 — 现有接口,向后兼容。"""
|
|
repo = request.app.state.repo
|
|
svc = BacktestService(repo)
|
|
end = req.end or date.today()
|
|
start = req.start or (end - timedelta(days=365 * 3))
|
|
|
|
cfg = BacktestConfig(
|
|
symbols=req.symbols,
|
|
start=start,
|
|
end=end,
|
|
entries=req.entries,
|
|
exits=req.exits,
|
|
stop_loss_pct=req.stop_loss_pct,
|
|
max_hold_days=req.max_hold_days,
|
|
fees_pct=req.fees_pct,
|
|
slippage_bps=req.slippage_bps,
|
|
matching=req.matching,
|
|
asset_type=req.asset_type,
|
|
)
|
|
try:
|
|
result = svc.run(cfg)
|
|
except VectorbtUnavailable as e:
|
|
raise HTTPException(status_code=503, detail=str(e)) from e
|
|
return asdict(result)
|
|
|
|
|
|
# ================================================================
|
|
# 因子回测
|
|
# ================================================================
|
|
|
|
class FactorColumnsResponse(BaseModel):
|
|
columns: list[dict]
|
|
|
|
|
|
@router.get("/factor/columns")
|
|
def factor_columns():
|
|
"""返回可用的因子列列表。"""
|
|
from app.backtest.factor import FACTOR_COLUMNS
|
|
return {"columns": FACTOR_COLUMNS}
|
|
|
|
|
|
class FactorBacktestRequest(BaseModel):
|
|
factor_name: str
|
|
symbols: list[str] | None = None
|
|
start: date | None = None
|
|
end: date | None = None
|
|
n_groups: int = 5
|
|
rebalance: Literal["daily", "weekly", "monthly"] = "monthly"
|
|
weight: Literal["equal", "factor_weight"] = "equal"
|
|
fees_pct: float = 0.0002
|
|
slippage_bps: float = 5.0
|
|
asset_type: str = "stock"
|
|
|
|
|
|
@router.post("/factor/run")
|
|
def factor_run(req: FactorBacktestRequest, request: Request):
|
|
"""因子回测 — IC/IR 分析 + 分层回测。"""
|
|
from app.backtest.factor import FactorBacktestService, FactorConfig
|
|
|
|
engine = _get_engine(request)
|
|
svc = FactorBacktestService(engine)
|
|
|
|
end = req.end or date.today()
|
|
start = _resolve_start(req, end, STRATEGY_DEFAULT_DAYS)
|
|
_guard_server_backtest_range(start, end)
|
|
symbols = req.symbols if req.symbols else None
|
|
if symbols is not None and len(symbols) > FACTOR_MAX_SYMBOLS:
|
|
raise HTTPException(
|
|
status_code=400,
|
|
detail=f"指定标的最多支持 {FACTOR_MAX_SYMBOLS} 只,请缩小标的范围。",
|
|
)
|
|
|
|
cfg = FactorConfig(
|
|
factor_name=req.factor_name,
|
|
symbols=symbols,
|
|
start=start,
|
|
end=end,
|
|
n_groups=req.n_groups,
|
|
rebalance=req.rebalance,
|
|
weight=req.weight,
|
|
fees_pct=req.fees_pct,
|
|
slippage_bps=req.slippage_bps,
|
|
asset_type=req.asset_type,
|
|
)
|
|
result = svc.run(cfg)
|
|
return asdict(result)
|
|
|
|
|
|
# ================================================================
|
|
# 策略回测
|
|
# ================================================================
|
|
|
|
class StrategyBacktestRequest(BaseModel):
|
|
strategy_id: str
|
|
symbols: list[str] | None = None
|
|
start: date | None = None
|
|
end: date | None = None
|
|
params: dict | None = None
|
|
overrides: dict | None = None
|
|
# matching 向后兼容; 显式传 entry_fill/exit_fill 时以二者为准。
|
|
matching: Literal["close_t", "open_t+1"] = "open_t+1"
|
|
entry_fill: Literal["close_t", "open_t+1"] | None = None
|
|
exit_fill: Literal["close_t", "open_t+1"] | None = None
|
|
fees_pct: float = 0.0002
|
|
commission_pct: float | None = None
|
|
stamp_tax_pct: float | None = None
|
|
slippage_bps: float = 5.0
|
|
max_positions: int = 10
|
|
max_exposure_pct: float = 1.0
|
|
initial_capital: float = 1_000_000.0
|
|
position_sizing: Literal["equal", "score_weight"] = "equal"
|
|
mode: Literal["position", "full"] = "position"
|
|
holding_days: int = 5
|
|
asset_type: str = "stock"
|
|
|
|
|
|
@router.post("/strategy/run")
|
|
def strategy_run(req: StrategyBacktestRequest, request: Request):
|
|
"""策略回测 — 复用 StrategyDef 体系做全周期回测。"""
|
|
from app.backtest.strategy import StrategyBacktestService, StrategyBacktestConfig
|
|
|
|
engine = _get_engine(request)
|
|
strategy_engine = request.app.state.strategy_engine
|
|
svc = StrategyBacktestService(engine, strategy_engine)
|
|
|
|
end = req.end or date.today()
|
|
start = _resolve_start(req, end, FACTOR_DEFAULT_DAYS)
|
|
_guard_server_backtest_range(start, end)
|
|
|
|
cfg = StrategyBacktestConfig(
|
|
strategy_id=req.strategy_id,
|
|
symbols=req.symbols if req.symbols else None,
|
|
start=start,
|
|
end=end,
|
|
params=req.params,
|
|
overrides=req.overrides,
|
|
matching=req.matching,
|
|
entry_fill=req.entry_fill,
|
|
exit_fill=req.exit_fill,
|
|
fees_pct=req.fees_pct,
|
|
commission_pct=req.commission_pct,
|
|
stamp_tax_pct=req.stamp_tax_pct,
|
|
slippage_bps=req.slippage_bps,
|
|
max_positions=req.max_positions,
|
|
max_exposure_pct=req.max_exposure_pct,
|
|
initial_capital=req.initial_capital,
|
|
position_sizing=req.position_sizing,
|
|
mode=req.mode,
|
|
holding_days=req.holding_days,
|
|
asset_type=req.asset_type,
|
|
)
|
|
result = svc.run(cfg)
|
|
return asdict(result)
|
|
|
|
|
|
# ── SSE 流式回测 (实时进度 + 可取消 + 支持重连) ───────────────────
|
|
|
|
import time
|
|
import hashlib
|
|
|
|
|
|
class _BacktestJob:
|
|
"""单个回测任务的状态, 存模块级供重连使用。"""
|
|
__slots__ = ("key", "cancel_event", "progress", "result", "error", "done", "finish_ts")
|
|
|
|
def __init__(self, key: str):
|
|
self.key = key
|
|
self.cancel_event = threading.Event()
|
|
self.progress: list[dict] = [] # 进度历史 (新连接可回放)
|
|
self.result = None # 完成后的结果
|
|
self.error: str | None = None
|
|
self.done = False
|
|
self.finish_ts: float = 0.0
|
|
|
|
|
|
# 模块级任务表: key -> _BacktestJob
|
|
_running_jobs: dict[str, _BacktestJob] = {}
|
|
_jobs_lock = threading.Lock()
|
|
_JOB_TTL = 300 # 完成后保留 5 分钟
|
|
|
|
# 并发回测上限: 多个重回测同时跑会 OOM (服务器内存约 1.8GB)。用信号量限并发,
|
|
# 超出的任务在 _run_backtest 里排队, SSE 连接照常保持, run 一开始就有进度。
|
|
_backtest_semaphore = threading.Semaphore(2)
|
|
|
|
|
|
def _cleanup_stale_jobs():
|
|
"""清理过期任务 (完成超过 TTL 的)。全程持 _jobs_lock: 迭代+pop 与其他访问互斥。"""
|
|
now = time.time()
|
|
with _jobs_lock:
|
|
stale = [k for k, j in _running_jobs.items() if j.done and now - j.finish_ts > _JOB_TTL]
|
|
for k in stale:
|
|
_running_jobs.pop(k, None)
|
|
|
|
|
|
def _make_job_key(
|
|
strategy_id: str, symbols: str | None, start: str | None, end: str | None,
|
|
matching: str, entry_fill: str | None, exit_fill: str | None,
|
|
fees_pct: float, slippage_bps: float,
|
|
max_positions: int, max_exposure_pct: float, initial_capital: float, position_sizing: str,
|
|
params: str | None, overrides: str | None,
|
|
mode: str = "position", holding_days: int = 5,
|
|
commission_pct: float | None = None, stamp_tax_pct: float | None = None,
|
|
asset_type: str = "stock",
|
|
) -> str:
|
|
raw = f"{strategy_id}|{symbols}|{start}|{end}|{matching}|{entry_fill}|{exit_fill}|{fees_pct}|{slippage_bps}|{max_positions}|{max_exposure_pct}|{initial_capital}|{position_sizing}|{params}|{overrides}|{mode}|{holding_days}|{commission_pct}|{stamp_tax_pct}|{asset_type}"
|
|
return hashlib.md5(raw.encode()).hexdigest()[:12]
|
|
|
|
|
|
@router.get("/strategy/stream")
|
|
async def strategy_stream(
|
|
request: Request,
|
|
strategy_id: str,
|
|
symbols: str | None = None,
|
|
start: str | None = None,
|
|
end: str | None = None,
|
|
matching: str = "open_t+1",
|
|
entry_fill: str | None = None,
|
|
exit_fill: str | None = None,
|
|
fees_pct: float = 0.0002,
|
|
commission_pct: float | None = None,
|
|
stamp_tax_pct: float | None = None,
|
|
slippage_bps: float = 5.0,
|
|
max_positions: int = 10,
|
|
max_exposure_pct: float = 1.0,
|
|
initial_capital: float = 1_000_000.0,
|
|
position_sizing: str = "equal",
|
|
params: str | None = None,
|
|
overrides: str | None = None,
|
|
mode: str = "position",
|
|
holding_days: int = 5,
|
|
asset_type: str = "stock",
|
|
):
|
|
"""SSE 流式策略回测: 实时推送进度, 完成后推送结果, 支持重连 (刷新/切页后恢复)。
|
|
|
|
- 相同参数的任务只启动一次, 多次连接订阅同一个任务
|
|
- 断开连接不会取消任务 (除非显式调用 cancel)
|
|
- 结果保留 5 分钟供重连
|
|
|
|
事件类型:
|
|
- progress: {day, total, date, equity}
|
|
- done: {result} (完整回测结果)
|
|
- error: {message}
|
|
"""
|
|
from app.backtest.strategy import StrategyBacktestService, StrategyBacktestConfig
|
|
|
|
engine = _get_engine(request)
|
|
strategy_engine = request.app.state.strategy_engine
|
|
svc = StrategyBacktestService(engine, strategy_engine)
|
|
|
|
end_date = date.fromisoformat(end) if end else date.today()
|
|
if start:
|
|
start_date = date.fromisoformat(start)
|
|
else:
|
|
# 空 start = 全部历史: 用本地最早日K日期, 查不到再回退到默认窗口
|
|
earliest = request.app.state.repo.earliest_daily_date()
|
|
start_date = earliest or (end_date - timedelta(days=FACTOR_DEFAULT_DAYS))
|
|
|
|
# 服务端范围保护
|
|
guard_violated = False
|
|
if settings.backtest_range_guard:
|
|
days = (end_date - start_date).days + 1
|
|
if days > BACKTEST_MAX_SERVER_DAYS:
|
|
guard_violated = True
|
|
|
|
job_key = _make_job_key(
|
|
strategy_id, symbols, start, end,
|
|
matching, entry_fill, exit_fill,
|
|
fees_pct, slippage_bps, max_positions, max_exposure_pct, initial_capital, position_sizing,
|
|
params, overrides,
|
|
mode, holding_days,
|
|
commission_pct, stamp_tax_pct,
|
|
asset_type=asset_type,
|
|
)
|
|
|
|
_cleanup_stale_jobs()
|
|
|
|
# 获取或创建任务
|
|
with _jobs_lock:
|
|
job = _running_jobs.get(job_key)
|
|
if job is None:
|
|
job = _BacktestJob(job_key)
|
|
_running_jobs[job_key] = job
|
|
is_new = True
|
|
else:
|
|
is_new = False
|
|
|
|
async def event_generator():
|
|
# 范围保护: 直接报错
|
|
if guard_violated:
|
|
yield f"event: error\ndata: {json.dumps({'message': BACKTEST_SERVER_GUARD_MESSAGE}, ensure_ascii=False)}\n\n"
|
|
return
|
|
|
|
# 如果是新任务, 启动回测线程
|
|
if is_new and not job.done:
|
|
cfg = StrategyBacktestConfig(
|
|
strategy_id=strategy_id,
|
|
symbols=[s.strip() for s in symbols.split(",") if s.strip()] if symbols else None,
|
|
start=start_date,
|
|
end=end_date,
|
|
params=json.loads(params) if params else None,
|
|
overrides=json.loads(overrides) if overrides else None,
|
|
matching=matching,
|
|
entry_fill=entry_fill,
|
|
exit_fill=exit_fill,
|
|
fees_pct=fees_pct,
|
|
commission_pct=commission_pct,
|
|
stamp_tax_pct=stamp_tax_pct,
|
|
slippage_bps=slippage_bps,
|
|
max_positions=int(max_positions),
|
|
max_exposure_pct=float(max_exposure_pct),
|
|
initial_capital=float(initial_capital),
|
|
position_sizing=position_sizing,
|
|
mode=mode,
|
|
holding_days=int(holding_days),
|
|
asset_type=asset_type,
|
|
)
|
|
|
|
def _run_backtest():
|
|
# 信号量限并发: 超额任务在此阻塞排队, 不并发吃满内存 (等待期间 cancel_event
|
|
# 仍可置位, svc.run 会据此提前返回 cancelled)。持槽跑完在 finally 释放。
|
|
_backtest_semaphore.acquire()
|
|
try:
|
|
result = svc.run(cfg, lambda d: job.progress.append(d), job.cancel_event)
|
|
job.result = result
|
|
job.done = True
|
|
job.finish_ts = time.time()
|
|
except Exception as e:
|
|
job.error = str(e)
|
|
job.done = True
|
|
job.finish_ts = time.time()
|
|
finally:
|
|
_backtest_semaphore.release()
|
|
|
|
# 启动后台线程 (不阻塞事件循环)
|
|
threading.Thread(target=_run_backtest, daemon=True).start()
|
|
|
|
# 订阅进度: 用读指针读 job.progress 列表 (多连接互不干扰)
|
|
cursor = 0
|
|
tick = 0
|
|
|
|
try:
|
|
while True:
|
|
# 已完成: 推送最终结果/错误并退出
|
|
if job.done:
|
|
if job.error:
|
|
yield f"event: error\ndata: {json.dumps({'message': job.error}, ensure_ascii=False)}\n\n"
|
|
elif job.result is not None:
|
|
r = job.result
|
|
if hasattr(r, "error") and r.error == "cancelled":
|
|
yield f"event: error\ndata: {json.dumps({'message': '回测已取消'}, ensure_ascii=False)}\n\n"
|
|
elif hasattr(r, "error") and r.error:
|
|
yield f"event: error\ndata: {json.dumps({'message': r.error}, ensure_ascii=False)}\n\n"
|
|
else:
|
|
yield f"event: done\ndata: {json.dumps(asdict(r), ensure_ascii=False, default=str)}\n\n"
|
|
return
|
|
|
|
# 断开检测: 每 4 轮检查一次 (降低 GIL 抢占频率)
|
|
tick += 1
|
|
if tick % 4 == 0 and await request.is_disconnected():
|
|
break
|
|
|
|
# 推送新进度 (从 cursor 开始读)
|
|
prog_list = job.progress
|
|
while cursor < len(prog_list):
|
|
msg = prog_list[cursor]
|
|
cursor += 1
|
|
yield f"event: progress\ndata: {json.dumps(msg, ensure_ascii=False, default=str)}\n\n"
|
|
|
|
await asyncio.sleep(0.5)
|
|
|
|
except asyncio.CancelledError:
|
|
raise
|
|
|
|
return StreamingResponse(event_generator(), media_type="text/event-stream")
|
|
|
|
|
|
@router.post("/strategy/cancel")
|
|
async def strategy_cancel(request: Request):
|
|
"""取消正在运行的回测任务 (前端传 query string, 后端算 job_key)。"""
|
|
body = await request.json()
|
|
qs = body.get("qs", "")
|
|
# 解析 qs 得到参数
|
|
from urllib.parse import parse_qs
|
|
p = parse_qs(qs)
|
|
def _get(key: str, default: str = "") -> str:
|
|
return p.get(key, [default])[0]
|
|
def _get_opt_float(key: str) -> float | None:
|
|
# 可选成本参数: 缺省或空串 → None (与 stream 侧 float | None 口径一致, 保证 job_key 对齐)。
|
|
v = _get(key)
|
|
return float(v) if v else None
|
|
job_key = _make_job_key(
|
|
_get("strategy_id"),
|
|
_get("symbols") or None,
|
|
_get("start") or None,
|
|
_get("end") or None,
|
|
_get("matching", "open_t+1"),
|
|
_get("entry_fill") or None,
|
|
_get("exit_fill") or None,
|
|
float(_get("fees_pct", "0.0002")),
|
|
float(_get("slippage_bps", "5")),
|
|
int(_get("max_positions", "10")),
|
|
float(_get("max_exposure_pct", "1")),
|
|
float(_get("initial_capital", "1000000")),
|
|
_get("position_sizing", "equal"),
|
|
_get("params") or None,
|
|
_get("overrides") or None,
|
|
_get("mode", "position"),
|
|
int(_get("holding_days", "5")),
|
|
commission_pct=_get_opt_float("commission_pct"),
|
|
stamp_tax_pct=_get_opt_float("stamp_tax_pct"),
|
|
asset_type=_get("asset_type", "stock"),
|
|
)
|
|
# 持锁读任务表: 与 _cleanup_stale_jobs 的 pop、stream 的写入互斥
|
|
with _jobs_lock:
|
|
job = _running_jobs.get(job_key)
|
|
if job and not job.done:
|
|
job.cancel_event.set()
|
|
return {"ok": True}
|
|
return {"ok": False, "message": "任务不存在或已完成"}
|
|
|