移除自选和回测功能

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-07-04 17:41:08 +08:00
parent ab6282b94c
commit a45aa5b033
49 changed files with 43 additions and 10091 deletions
-392
View File
@@ -1,392 +0,0 @@
"""回测服务(§6.7)。
包 vectorbt — 全项目唯一一处出现 pandas。
"""
from __future__ import annotations
import logging
import uuid
from dataclasses import dataclass, field
from datetime import date
from typing import Literal
import numpy as np
import pandas as pd
import polars as pl
from app.config import settings
from app.tickflow.repository import KlineRepository
logger = logging.getLogger(__name__)
# vectorbt 是 optional extras(见 pyproject.toml).未装时只有 backtest 不可用,其他功能正常.
_vbt = None
_vbt_unavailable_reason: str | None = None
class VectorbtUnavailable(RuntimeError):
"""vectorbt 未安装 — 提示用户 `uv sync --extra backtest`."""
def _get_vbt():
global _vbt, _vbt_unavailable_reason
if _vbt is not None:
return _vbt
if _vbt_unavailable_reason is not None:
raise VectorbtUnavailable(_vbt_unavailable_reason)
try:
import vectorbt as vbt
_vbt = vbt
return _vbt
except ImportError as e:
_vbt_unavailable_reason = (
"vectorbt 未安装 — 它是回测的可选依赖.macOS Intel 用户先 `brew install cmake` "
"然后 `uv sync --extra backtest`"
)
logger.warning("vectorbt unavailable: %s", e)
raise VectorbtUnavailable(_vbt_unavailable_reason) from e
def is_available() -> bool:
"""供 API 层快速检测."""
try:
_get_vbt()
return True
except VectorbtUnavailable:
return False
SignalKind = Literal[
"macd_golden", "macd_dead",
"ma_golden_5_20", "ma_dead_5_20",
"ma_golden_20_60",
"ma20_breakout", "ma20_breakdown",
"n_day_high", "n_day_low",
"boll_breakout_upper", "boll_breakdown_lower",
"volume_surge",
"rsi_oversold", "rsi_overbought",
"stop_loss", "trailing_stop", "max_hold",
]
@dataclass
class BacktestConfig:
symbols: list[str]
start: date
end: date
# 买入信号(任一触发即买)
entries: list[str] = field(default_factory=list)
# 卖出信号(任一触发即卖)
exits: list[str] = field(default_factory=list)
# 其他参数
stop_loss_pct: float | None = None # 例 -0.05 = -5%
max_hold_days: int | None = None
fees_pct: float = 0.0002 # 万二佣金
slippage_bps: float = 5 # 5 bps
# 撮合
matching: Literal["close_t", "open_t+1"] = "close_t"
rsi_oversold_threshold: float = 30
rsi_overbought_threshold: float = 70
@dataclass
class BacktestResult:
run_id: str
config: dict
stats: dict
equity_curve: list[dict] # [{date, value}]
trades: list[dict] # [{symbol, entry_date, exit_date, pnl_pct, ...}]
per_symbol_stats: list[dict] # 每只股票的统计
# enriched 表里的信号列名映射
_SIGNAL_COLS: dict[SignalKind, str] = {
"macd_golden": "signal_macd_golden",
"macd_dead": "signal_macd_dead",
"ma_golden_5_20": "signal_ma_golden_5_20",
"ma_dead_5_20": "signal_ma_dead_5_20",
"ma_golden_20_60": "signal_ma_golden_20_60",
"ma20_breakout": "signal_ma20_breakout",
"ma20_breakdown": "signal_ma20_breakdown",
"n_day_high": "signal_n_day_high",
"n_day_low": "signal_n_day_low",
"boll_breakout_upper": "signal_boll_breakout_upper",
"boll_breakdown_lower": "signal_boll_breakdown_lower",
"volume_surge": "signal_volume_surge",
}
class BacktestService:
def __init__(self, repo: KlineRepository) -> None:
self.repo = repo
def _load_panel(
self,
symbols: list[str],
start: date,
end: date,
) -> pd.DataFrame:
"""加载 [date × symbol] 价格面板 — Polars scan_parquet + 即时计算指标。
**全项目唯一从 Polars 转 pandas 的边界**(§7.4 / ADR-19)。
"""
try:
enriched_glob = str(self.repo.store.data_dir / "kline_daily_enriched" / "**" / "*.parquet")
df = (
pl.scan_parquet(enriched_glob)
.filter(
(pl.col("symbol").is_in(symbols))
& (pl.col("date") >= start)
& (pl.col("date") <= end)
)
.sort(["date", "symbol"])
.collect()
)
except Exception as e: # noqa: BLE001
logger.warning("backtest load failed: %s", e)
return pd.DataFrame()
if df.is_empty():
return pd.DataFrame()
# 即时计算指标 + 信号
from app.indicators.pipeline import compute_all
df = compute_all(df)
# 选择需要的列
needed_cols = [
"date", "symbol", "open", "high", "low", "close", "volume",
"rsi_14", "signal_macd_golden", "signal_macd_dead",
"signal_ma_golden_5_20", "signal_ma_dead_5_20",
"signal_ma_golden_20_60",
"signal_ma20_breakout", "signal_ma20_breakdown",
"signal_n_day_high", "signal_n_day_low",
"signal_boll_breakout_upper", "signal_boll_breakdown_lower",
"signal_volume_surge",
]
existing = [c for c in needed_cols if c in df.columns]
df = df.select(existing)
# to_pandas 边界
return df.to_pandas(use_pyarrow_extension_array=False)
def _build_signal_matrix(
self,
panel: pd.DataFrame,
kinds: list[str],
config: BacktestConfig,
) -> pd.DataFrame:
"""从面板构造 [date × symbol] 的布尔信号矩阵。"""
if not kinds or panel.empty:
return pd.DataFrame()
# pivot 成 [date × symbol] 形式
result = None
for kind in kinds:
mat = None
if kind in _SIGNAL_COLS:
col = _SIGNAL_COLS[kind]
mat = panel.pivot(index="date", columns="symbol", values=col).fillna(False).astype(bool)
elif kind == "rsi_oversold":
mat = (panel.pivot(index="date", columns="symbol", values="rsi_14")
< config.rsi_oversold_threshold)
elif kind == "rsi_overbought":
mat = (panel.pivot(index="date", columns="symbol", values="rsi_14")
> config.rsi_overbought_threshold)
# stop_loss / trailing / max_hold 通过 vectorbt 参数处理,不参与信号矩阵
if mat is not None:
result = mat if result is None else (result | mat)
return result if result is not None else pd.DataFrame()
def run(self, config: BacktestConfig) -> BacktestResult:
vbt = _get_vbt()
run_id = uuid.uuid4().hex[:10]
panel = self._load_panel(config.symbols, config.start, config.end)
if panel.empty:
return BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={"error": "no data"},
equity_curve=[],
trades=[],
per_symbol_stats=[],
)
# 价格面板
close = panel.pivot(index="date", columns="symbol", values="close")
# 信号矩阵
entries = self._build_signal_matrix(panel, config.entries, config)
exits = self._build_signal_matrix(panel, config.exits, config)
# 对齐 index/columns
if not entries.empty:
entries = entries.reindex_like(close).fillna(False).astype(bool)
else:
entries = pd.DataFrame(False, index=close.index, columns=close.columns)
if not exits.empty:
exits = exits.reindex_like(close).fillna(False).astype(bool)
else:
exits = pd.DataFrame(False, index=close.index, columns=close.columns)
if not entries.any().any():
return BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={"error": "no buy signals"},
equity_curve=[],
trades=[],
per_symbol_stats=[],
)
# T+1 适配:vectorbt 默认信号当根 K 撮合
# close_t 撮合:维持默认
# open_t+1 撮合:shift 信号 1 根 + 用 open 作为价
if config.matching == "open_t+1":
entries = entries.shift(1).fillna(False).astype(bool)
exits = exits.shift(1).fillna(False).astype(bool)
price = panel.pivot(index="date", columns="symbol", values="open")
else:
price = close
# 跑回测
try:
pf_kwargs = dict(
close=close,
entries=entries,
exits=exits,
price=price,
fees=config.fees_pct,
slippage=config.slippage_bps / 10000.0,
freq="1D",
)
if config.stop_loss_pct is not None:
pf_kwargs["sl_stop"] = abs(config.stop_loss_pct)
if config.max_hold_days is not None:
# vectorbt 没有内置 max-hold;用时间退出近似:
# 在 max_hold_days 后强制 exit
exits_idx = entries.copy()
for col in entries.columns:
entry_rows = np.where(entries[col].values)[0]
for i in entry_rows:
end_i = min(i + config.max_hold_days, len(entries) - 1)
if end_i > i:
exits_idx.iloc[end_i][col] = True
pf_kwargs["exits"] = (exits | exits_idx).astype(bool)
pf = vbt.Portfolio.from_signals(**pf_kwargs)
except Exception as e: # noqa: BLE001
logger.exception("vectorbt backtest failed")
return BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={"error": str(e)},
equity_curve=[],
trades=[],
per_symbol_stats=[],
)
# 提取结果
try:
stats_series = pf.stats(silence_warnings=True)
if isinstance(stats_series, pd.DataFrame):
# 多列时取 agg
stats_dict = stats_series.mean(numeric_only=True).to_dict()
else:
stats_dict = stats_series.to_dict()
except Exception: # noqa: BLE001
stats_dict = {}
# 净值曲线(组合平均)
equity = pf.value().mean(axis=1) if isinstance(pf.value(), pd.DataFrame) else pf.value()
equity_curve = [
{"date": str(idx.date() if hasattr(idx, "date") else idx), "value": float(v)}
for idx, v in equity.items() if pd.notna(v)
]
# 交易记录
try:
trades_df = pf.trades.records_readable
trades = trades_df.to_dict(orient="records") if not trades_df.empty else []
# 字段名美化
trades = [
{
"symbol": t.get("Column", t.get("Symbol", "")),
"entry_date": str(t.get("Entry Timestamp", t.get("Entry Date", ""))),
"exit_date": str(t.get("Exit Timestamp", t.get("Exit Date", ""))),
"entry_price": float(t.get("Avg Entry Price", t.get("Avg. Entry Price", 0))),
"exit_price": float(t.get("Avg Exit Price", t.get("Avg. Exit Price", 0))),
"pnl_pct": float(t.get("Return", t.get("PnL %", 0))),
"duration": str(t.get("Duration", "")),
}
for t in trades
]
except Exception: # noqa: BLE001
trades = []
# 每标的统计
per_symbol = []
try:
total_ret = pf.total_return()
if isinstance(total_ret, pd.Series):
for sym, ret in total_ret.items():
if pd.notna(ret):
per_symbol.append({"symbol": sym, "total_return": float(ret)})
except Exception: # noqa: BLE001
pass
result = BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={k: _json_safe(v) for k, v in stats_dict.items()},
equity_curve=equity_curve,
trades=trades,
per_symbol_stats=per_symbol,
)
# 落盘
self._persist(result)
return result
def _persist(self, result: BacktestResult) -> None:
out_dir = settings.data_dir / "backtest_results"
out_dir.mkdir(parents=True, exist_ok=True)
# 用 polars 写一份汇总
summary = pl.DataFrame({
"run_id": [result.run_id],
"stats_json": [str(result.stats)],
"n_trades": [len(result.trades)],
})
summary.write_parquet(out_dir / f"run_id={result.run_id}.parquet")
def get_result(self, run_id: str) -> BacktestResult | None:
# Phase 1:只保留近似落盘,完整结果保存在内存的近期 cache 中
# 简化:重新 run 比缓存复杂结果代价小,暂不实现 get_result
return None
def _config_to_dict(c: BacktestConfig) -> dict:
return {
"symbols": c.symbols,
"start": str(c.start),
"end": str(c.end),
"entries": c.entries,
"exits": c.exits,
"stop_loss_pct": c.stop_loss_pct,
"max_hold_days": c.max_hold_days,
"fees_pct": c.fees_pct,
"slippage_bps": c.slippage_bps,
"matching": c.matching,
}
def _json_safe(v):
if isinstance(v, (int, float, str, bool)) or v is None:
return v
if isinstance(v, (np.floating, np.integer)):
return float(v) if not np.isnan(float(v)) else None
if hasattr(v, "isoformat"):
return v.isoformat()
return str(v)
+2 -3
View File
@@ -38,7 +38,7 @@ def _invalidate(table: str | None = None) -> None:
def _resolve_universe(capset: CapabilitySet) -> list[str]:
"""解析标的池 — 与 daily_pipeline 独立的副本"""
"""解析标的池。"""
if capset.has(Cap.KLINE_DAILY_BATCH):
try:
from app.tickflow.pools import get_pool
@@ -48,12 +48,11 @@ def _resolve_universe(capset: CapabilitySet) -> list[str]:
except Exception as e:
logger.warning("CN_Equity_A pool unavailable: %s", e)
from app.tickflow.pools import DEMO_SYMBOLS, get_pool as _get_pool
from app.tickflow.pools import DEMO_SYMBOLS
from app.config import settings
from pathlib import Path
import polars as pl
base: set[str] = set(DEMO_SYMBOLS)
base.update(_get_pool("watchlist"))
d = Path(settings.data_dir)
inst_path = d / "instruments" / "instruments.parquet"
if inst_path.exists():
-11
View File
@@ -296,17 +296,6 @@ def set_nav_hidden(hidden: list[str]) -> list[str]:
return get_nav_hidden()
def get_watchlist_columns() -> list[dict] | None:
"""返回自选列表列配置。"""
return load().get("watchlist_columns")
def set_watchlist_columns(columns: list[dict]) -> list[dict]:
"""保存自选列表列配置。"""
save({"watchlist_columns": columns})
return columns
def get_screener_result_columns() -> list[dict] | None:
"""返回策略结果列表列配置。"""
return load().get("screener_result_columns")
-144
View File
@@ -1,144 +0,0 @@
"""自选股服务(§6.1)。
存储:`data/user_data/watchlist.parquet`,字段 symbol + added_at + note。
"""
from __future__ import annotations
import logging
from datetime import datetime
from pathlib import Path
import polars as pl
from app.config import settings
from app.tickflow.capabilities import Cap, CapabilitySet
from app.tickflow.client import get_client
logger = logging.getLogger(__name__)
def _path() -> Path:
p = settings.data_dir / "user_data" / "watchlist.parquet"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def list_symbols() -> list[dict]:
p = _path()
if not p.exists():
return []
df = pl.read_parquet(p)
if df.is_empty():
return []
return df.to_dicts()
def add(symbol: str, note: str = "") -> list[dict]:
p = _path()
if p.exists():
df = pl.read_parquet(p)
# 已存在则先移除,后面重新插入到最前面
if symbol in df["symbol"].to_list():
df = df.filter(pl.col("symbol") != symbol)
else:
df = pl.DataFrame(schema={"symbol": pl.Utf8, "added_at": pl.Utf8, "note": pl.Utf8})
new_row = pl.DataFrame({
"symbol": [symbol],
"added_at": [datetime.utcnow().isoformat(timespec="seconds")],
"note": [note],
})
out = pl.concat([new_row, df], how="diagonal_relaxed")
out.write_parquet(p)
return out.to_dicts()
def remove(symbol: str) -> list[dict]:
p = _path()
if not p.exists():
return []
df = pl.read_parquet(p)
df = df.filter(pl.col("symbol") != symbol)
df.write_parquet(p)
return df.to_dicts()
def move_to_top(symbol: str) -> list[dict]:
p = _path()
if not p.exists():
return []
df = pl.read_parquet(p)
if df.is_empty() or symbol not in df["symbol"].to_list():
return df.to_dicts()
target = df.filter(pl.col("symbol") == symbol)
rest = df.filter(pl.col("symbol") != symbol)
out = pl.concat([target, rest], how="diagonal_relaxed")
out.write_parquet(p)
return out.to_dicts()
def clear() -> int:
"""清空自选列表。返回移除的数量。"""
p = _path()
if not p.exists():
return 0
df = pl.read_parquet(p)
count = df.height
if count > 0:
pl.DataFrame(schema={"symbol": pl.Utf8, "added_at": pl.Utf8, "note": pl.Utf8}).write_parquet(p)
return count
def fetch_quotes(symbols: list[str], capset: CapabilitySet, timeout_s: float = 8.0) -> list[dict]:
"""拉取实时行情。
优先用 quote.batch;否则降级为 quote.by_symbol 单股请求。
timeout_s: 单批次请求超时(秒),防止 API 卡死阻塞整个请求。
"""
from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeout
if not symbols:
return []
tf = get_client()
quotes: list[dict] = []
# 走 batch
batch_size = 5
if capset.has(Cap.QUOTE_BATCH):
lim = capset.limits(Cap.QUOTE_BATCH)
batch_size = lim.batch if lim and lim.batch else 50
elif capset.has(Cap.QUOTE_BY_SYMBOL):
lim = capset.limits(Cap.QUOTE_BY_SYMBOL)
batch_size = lim.batch if lim and lim.batch else 5
else:
# 无任何实时行情能力(none/free 档走 free-api 服务器,不提供实时行情)
# 提前返回空,避免发起注定失败的请求
return []
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
# 用线程池为每个批次加超时保护
pool = ThreadPoolExecutor(max_workers=1)
for chunk in chunks:
try:
future = pool.submit(tf.quotes.get, symbols=chunk, as_dataframe=True)
raw = future.result(timeout=timeout_s)
if raw is None or len(raw) == 0:
continue
df = pl.from_pandas(raw)
rename_map = {
"last_price": "price",
"ext.change_pct": "pct",
"ext.name": "name",
}
df = df.rename({k: v for k, v in rename_map.items() if k in df.columns})
quotes.extend(df.to_dicts())
except FuturesTimeout:
logger.warning("quote fetch timeout (%.1fs) for %d symbols", timeout_s, len(chunk))
break # 超时后不再尝试后续批次
except Exception as e: # noqa: BLE001
logger.warning("quote fetch failed for %d symbols: %s", len(chunk), e)
pool.shutdown(wait=False)
return quotes