Add walk-forward backtest optimization to mitigate signal overfitting (item m)
Rolling train/test folds over 3 years of data auto-optimize the three cheap-to-tune trading parameters (STRONG/BUY score thresholds, max hold time) via grid search on each fold's train window, then evaluate purely on the held-out test window. Stitching all out-of-sample results gives an honest performance estimate uninflated by tuning against the same data used to score it. Split signal_scoring.py's expensive 13-algorithm scoring from its cheap final threshold classification so grid search can replay many parameter combinations without recomputing indicators each time. Moved the backtest engine (fetch/precompute/simulate) out of the API layer into app/services/backtest_engine.py so both /backtest/run and the new walk-forward optimizer share one implementation instead of drifting copies — same rationale as the earlier signal_service.py split (item h). Also merges two long-diverged Alembic migration heads discovered while adding the walk_forward_results table, so `alembic upgrade head` has a single target again. New: POST/GET/DELETE /walk-forward/* endpoints, a Walk-Forward tab on the Backtest page (fold table, out-of-sample equity curve, run history). 19 new backend tests (153 total, all passing). Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -1,10 +1,7 @@
|
||||
"""Backtest API endpoint — run backtest and return JSON results."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from collections import defaultdict
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import select, and_, func
|
||||
@@ -14,21 +11,11 @@ from app.database import get_db
|
||||
from app.models.candle import Candle
|
||||
from app.models.symbol import Symbol
|
||||
from app.models.exchange import Exchange
|
||||
from app.services.indicator_service import (
|
||||
bollinger_bands, rsi, sma, macd, supertrend,
|
||||
volume_breakout, ichimoku, detect_divergence, market_structure,
|
||||
)
|
||||
from app.services.signal_service import (
|
||||
_classify_signal_combined,
|
||||
STRONG_BUY, BUY, STRONG_SELL, SELL,
|
||||
)
|
||||
from app.services.backtest_engine import run_backtest as _run_backtest, MIN_CANDLES
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/backtest", tags=["backtest"])
|
||||
|
||||
TRADE_SIZE = Decimal("10")
|
||||
MAX_HOLD_CANDLES = 48
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def backtest_root():
|
||||
@@ -36,315 +23,12 @@ async def backtest_root():
|
||||
return {"message": "Use GET /backtest/run to run a backtest"}
|
||||
|
||||
|
||||
async def _run_backtest(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
exchange: str,
|
||||
timeframe: str = "30m",
|
||||
days: int = 7,
|
||||
trade_size: Decimal = Decimal("10"),
|
||||
) -> dict:
|
||||
"""Run backtest and return structured results."""
|
||||
result = await db.execute(
|
||||
select(Symbol)
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(and_(Exchange.name == exchange, Symbol.symbol == symbol))
|
||||
)
|
||||
db_symbol = result.scalar_one_or_none()
|
||||
if not db_symbol:
|
||||
return {"error": f"Symbol {symbol} not found on {exchange}"}
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
|
||||
result = await db.execute(
|
||||
select(Candle)
|
||||
.where(and_(
|
||||
Candle.symbol_id == db_symbol.id,
|
||||
Candle.timeframe == timeframe,
|
||||
Candle.timestamp >= cutoff,
|
||||
))
|
||||
.order_by(Candle.timestamp.asc())
|
||||
)
|
||||
candles = list(result.scalars().all())
|
||||
|
||||
# Min candles for warmup: BB(20) + RSI(14) + some room = 30
|
||||
MIN_CANDLES = 30
|
||||
if len(candles) < MIN_CANDLES:
|
||||
if candles:
|
||||
span_hours = (candles[-1].timestamp - candles[0].timestamp).total_seconds() / 3600
|
||||
if span_hours >= 24:
|
||||
avail = f"{span_hours/24:.0f}d"
|
||||
else:
|
||||
avail = f"{span_hours:.0f}h"
|
||||
return {"error": f"Only ~{avail} data available (need at least {MIN_CANDLES} candles for {timeframe}). Try more days or a higher timeframe (4h)."}
|
||||
return {"error": f"No candle data found for {symbol} on {timeframe}. The exchange may not support this pair."}
|
||||
|
||||
# MTF config
|
||||
tf_minutes = {"15m": 15, "30m": 30, "1h": 60, "4h": 240}
|
||||
main_minutes = tf_minutes.get(timeframe, 30)
|
||||
mtf_config = []
|
||||
for mtf_tf, mtf_minutes, mtf_w in [("15m", 15, 0.5), ("1h", 60, 1.5), ("4h", 240, 2.0)]:
|
||||
if mtf_tf == timeframe:
|
||||
continue
|
||||
mult = mtf_minutes // main_minutes
|
||||
if mult >= 1 and len(candles) >= mult * MIN_CANDLES:
|
||||
mtf_config.append((mtf_tf, mult, mtf_w))
|
||||
|
||||
# ── Pre-compute all candle data & indicators ONCE ──
|
||||
candle_dicts_full = [
|
||||
{"high": float(c.high), "low": float(c.low),
|
||||
"close": float(c.close), "open": float(c.open),
|
||||
"volume": float(c.volume)}
|
||||
for c in candles
|
||||
]
|
||||
close_prices_full = [float(c.close) for c in candles]
|
||||
|
||||
# Pre-compute indicators on full dataset (O(n) instead of O(n²))
|
||||
bb_full = bollinger_bands(close_prices_full) or {}
|
||||
rsi_full = rsi(close_prices_full) or []
|
||||
sma_full = sma(close_prices_full, 20) or []
|
||||
macd_full = macd(close_prices_full) or {}
|
||||
st_full = supertrend(candle_dicts_full) or {}
|
||||
vb_full = volume_breakout(candle_dicts_full) or []
|
||||
ichi_full = ichimoku(candle_dicts_full) or {}
|
||||
smc_full = market_structure(candle_dicts_full) or {}
|
||||
|
||||
# Pre-compute divergence ONCE (uses full arrays, indexes match)
|
||||
rsi_div_full = detect_divergence(close_prices_full, rsi_full)
|
||||
macd_hist_full = macd_full.get("histogram", []) if macd_full else []
|
||||
macd_div_full = detect_divergence(close_prices_full, macd_hist_full)
|
||||
|
||||
# Pre-build MTF candles ONCE per MTF config
|
||||
mtf_precomputed = []
|
||||
for mtf_name, mtf_mult, mtf_w in mtf_config:
|
||||
mtf_candles_list = []
|
||||
for j in range(0, len(candle_dicts_full) - mtf_mult + 1, mtf_mult):
|
||||
chunk = candle_dicts_full[j:j + mtf_mult]
|
||||
mtf_candles_list.append({
|
||||
"open": chunk[0]["open"],
|
||||
"high": max(c["high"] for c in chunk),
|
||||
"low": min(c["low"] for c in chunk),
|
||||
"close": chunk[-1]["close"],
|
||||
"volume": sum(c["volume"] for c in chunk),
|
||||
})
|
||||
if len(mtf_candles_list) >= MIN_CANDLES:
|
||||
mtf_p = [c["close"] for c in mtf_candles_list]
|
||||
mtf_precomputed.append({
|
||||
"name": mtf_name,
|
||||
"weight": mtf_w,
|
||||
"mult": mtf_mult,
|
||||
"candles_list": mtf_candles_list,
|
||||
"close_prices": mtf_p,
|
||||
"bb": bollinger_bands(mtf_p) or {},
|
||||
"rsi": rsi(mtf_p) or [],
|
||||
"sma": sma(mtf_p, 20) or [],
|
||||
"macd": macd(mtf_p) or {},
|
||||
"st": supertrend(mtf_candles_list) or {},
|
||||
"vb": volume_breakout(mtf_candles_list) or [],
|
||||
"ichi": ichimoku(mtf_candles_list) or {},
|
||||
"smc": market_structure(mtf_candles_list) or {},
|
||||
})
|
||||
|
||||
all_signals = []
|
||||
trades = []
|
||||
current_position = None
|
||||
|
||||
for i in range(MIN_CANDLES, len(candles)):
|
||||
candle = candles[i]
|
||||
latest_close = close_prices_full[i]
|
||||
timestamp = candle.timestamp.isoformat()
|
||||
|
||||
# Slice pre-computed arrays (O(i) but ~100x faster than recomputing)
|
||||
clip = i + 1
|
||||
def _safe_slice(v):
|
||||
return v[:clip] if v is not None and hasattr(v, '__getitem__') else v
|
||||
bb_data = {k: _safe_slice(v) for k, v in bb_full.items()} if bb_full else {}
|
||||
rsi_data = rsi_full[:clip] if rsi_full else []
|
||||
sma_data = sma_full[:clip] if sma_full else []
|
||||
macd_data = {k: _safe_slice(v) for k, v in macd_full.items()} if macd_full else {}
|
||||
st_data = {k: _safe_slice(v) for k, v in st_full.items()} if st_full else {}
|
||||
vb_data = vb_full[:clip] if vb_full else []
|
||||
ichi_data = {k: _safe_slice(v) for k, v in ichi_full.items()} if ichi_full else {}
|
||||
smc_data = {k: _safe_slice(v) for k, v in smc_full.items()} if smc_full else {}
|
||||
|
||||
# MTF votes — use precomputed MTF indicators, sliced to current MTF candle index
|
||||
mtf_votes = []
|
||||
for mtf in mtf_precomputed:
|
||||
# Which MTF candle corresponds to main candle i?
|
||||
mtf_idx = i // mtf["mult"]
|
||||
if mtf_idx < MIN_CANDLES or mtf_idx >= len(mtf["close_prices"]):
|
||||
continue
|
||||
clip_mtf = mtf_idx + 1
|
||||
def _safe_slice_mtf(v):
|
||||
return v[:clip_mtf] if v is not None and hasattr(v, '__getitem__') else v
|
||||
mtf_s, *_ = _classify_signal_combined(
|
||||
mtf["close_prices"][mtf_idx],
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["bb"].items()},
|
||||
mtf["rsi"][:clip_mtf],
|
||||
mtf["sma"][:clip_mtf],
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["macd"].items()} if mtf["macd"] else None,
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["st"].items()} if mtf["st"] else None,
|
||||
mtf["vb"][:clip_mtf] if mtf["vb"] else None,
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["ichi"].items()} if mtf["ichi"] else None,
|
||||
(None, None), (None, None),
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["smc"].items()} if mtf["smc"] else None,
|
||||
)
|
||||
if mtf_s:
|
||||
mtf_votes.append((mtf_s, "", mtf["weight"]))
|
||||
|
||||
signal_type, strength, *_ = _classify_signal_combined(
|
||||
latest_close, bb_data, rsi_data, sma_data,
|
||||
macd_data, st_data, vb_data, ichi_data,
|
||||
rsi_div_full, macd_div_full, smc_data, mtf_votes or None,
|
||||
)
|
||||
|
||||
if signal_type:
|
||||
all_signals.append({
|
||||
"time": timestamp, "signal": signal_type,
|
||||
"strength": strength or "", "price": latest_close,
|
||||
})
|
||||
|
||||
# PnL simulation
|
||||
if signal_type in (STRONG_BUY, BUY):
|
||||
if current_position and current_position["direction"] == "SHORT":
|
||||
if signal_type == STRONG_BUY:
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
pnl = (entry - latest_close) * qty
|
||||
current_position.update({
|
||||
"exit_price": latest_close, "exit_time": timestamp,
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
|
||||
})
|
||||
trades.append(current_position)
|
||||
current_position = None
|
||||
else:
|
||||
continue
|
||||
if not current_position:
|
||||
qty = float(trade_size) / latest_close
|
||||
current_position = {
|
||||
"direction": "LONG", "entry_price": latest_close,
|
||||
"entry_time": timestamp, "quantity": qty,
|
||||
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
|
||||
}
|
||||
|
||||
elif signal_type in (STRONG_SELL, SELL):
|
||||
if current_position and current_position["direction"] == "LONG":
|
||||
if signal_type == STRONG_SELL:
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
pnl = (latest_close - entry) * qty
|
||||
current_position.update({
|
||||
"exit_price": latest_close, "exit_time": timestamp,
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
|
||||
})
|
||||
trades.append(current_position)
|
||||
current_position = None
|
||||
else:
|
||||
continue
|
||||
if not current_position:
|
||||
qty = float(trade_size) / latest_close
|
||||
current_position = {
|
||||
"direction": "SHORT", "entry_price": latest_close,
|
||||
"entry_time": timestamp, "quantity": qty,
|
||||
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
|
||||
}
|
||||
|
||||
# Time limit
|
||||
if current_position and current_position["status"] == "OPEN":
|
||||
hold = i - current_position["entry_index"]
|
||||
if hold >= MAX_HOLD_CANDLES:
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
if current_position["direction"] == "LONG":
|
||||
pnl = (latest_close - entry) * qty
|
||||
else:
|
||||
pnl = (entry - latest_close) * qty
|
||||
current_position.update({
|
||||
"exit_price": latest_close, "exit_time": timestamp,
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "TIME_LIMIT",
|
||||
})
|
||||
trades.append(current_position)
|
||||
current_position = None
|
||||
|
||||
# Close final position
|
||||
if current_position and current_position["status"] == "OPEN":
|
||||
last_close = float(candles[-1].close)
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
if current_position["direction"] == "LONG":
|
||||
pnl = (last_close - entry) * qty
|
||||
else:
|
||||
pnl = (entry - last_close) * qty
|
||||
current_position.update({
|
||||
"exit_price": last_close,
|
||||
"exit_time": candles[-1].timestamp.isoformat(),
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "END_OF_DATA",
|
||||
})
|
||||
trades.append(current_position)
|
||||
|
||||
# Compute stats
|
||||
counts = defaultdict(int)
|
||||
for s in all_signals:
|
||||
counts[s["signal"]] += 1
|
||||
|
||||
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
|
||||
winning_trades = [t for t in closed_trades if t.get("pnl", 0) > 0]
|
||||
losing_trades = [t for t in closed_trades if t.get("pnl", 0) <= 0]
|
||||
total_pnl = sum(t.get("pnl", 0) for t in closed_trades)
|
||||
gross_profit = sum(t.get("pnl", 0) for t in winning_trades)
|
||||
gross_loss = sum(t.get("pnl", 0) for t in losing_trades)
|
||||
win_rate = round(len(winning_trades) / len(closed_trades) * 100, 1) if closed_trades else 0
|
||||
profit_factor = round(abs(gross_profit / gross_loss), 2) if gross_loss != 0 else None
|
||||
avg_win = round(gross_profit / len(winning_trades), 2) if winning_trades else None
|
||||
avg_loss = round(gross_loss / len(losing_trades), 2) if losing_trades else None
|
||||
|
||||
best_trade = max(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
|
||||
worst_trade = min(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
|
||||
|
||||
return {
|
||||
"symbol": symbol,
|
||||
"exchange": exchange,
|
||||
"timeframe": timeframe,
|
||||
"days": days,
|
||||
"candles_count": len(candles),
|
||||
"signal_counts": dict(counts),
|
||||
"total_signals": len(all_signals),
|
||||
"recent_signals": all_signals[-15:],
|
||||
"trades": {
|
||||
"total": len(closed_trades),
|
||||
"wins": len(winning_trades),
|
||||
"losses": len(losing_trades),
|
||||
"win_rate": win_rate,
|
||||
"total_pnl": round(total_pnl, 2),
|
||||
"profit_factor": profit_factor,
|
||||
"avg_win": avg_win,
|
||||
"avg_loss": avg_loss,
|
||||
"best_trade": {
|
||||
"direction": best_trade.get("direction"),
|
||||
"entry_price": round(best_trade["entry_price"], 4),
|
||||
"exit_price": round(best_trade["exit_price"], 4),
|
||||
"pnl": round(best_trade["pnl"], 2),
|
||||
"entry_signal": best_trade.get("entry_signal"),
|
||||
} if best_trade else None,
|
||||
"worst_trade": {
|
||||
"direction": worst_trade.get("direction"),
|
||||
"entry_price": round(worst_trade["entry_price"], 4),
|
||||
"exit_price": round(worst_trade["exit_price"], 4),
|
||||
"pnl": round(worst_trade["pnl"], 2),
|
||||
"entry_signal": worst_trade.get("entry_signal"),
|
||||
} if worst_trade else None,
|
||||
"per_signal": {},
|
||||
"recent": closed_trades[-10:],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@router.get("/symbols")
|
||||
async def get_backtest_symbols(
|
||||
exchange: str = Query(None, description="Exchange name filter (e.g., binance, bybit)"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
"""Return symbols with sufficient candles (>=30 in each of 30m/1h/4h/1d) for backtesting."""
|
||||
MIN_CANDLES = 30
|
||||
TFS = ["30m", "1h", "4h", "1d"]
|
||||
|
||||
# Subquery: symbol_id + timeframe that have >= MIN_CANDLES
|
||||
|
||||
@@ -10,6 +10,7 @@ from app.api.v1.symbols import router as symbols_router
|
||||
from app.api.v1.signals import router as signals_router
|
||||
from app.api.v1.backtest import router as backtest_router
|
||||
from app.api.v1.backtest_history import router as backtest_history_router
|
||||
from app.api.v1.walk_forward import router as walk_forward_router
|
||||
from app.api.v1.watchlist import router as watchlist_router
|
||||
from app.api.v1.orders import router as orders_router
|
||||
from app.api.v1.real_trades import router as real_trades_router
|
||||
@@ -27,6 +28,7 @@ api_router.include_router(credentials_router)
|
||||
api_router.include_router(signals_router)
|
||||
api_router.include_router(backtest_router)
|
||||
api_router.include_router(backtest_history_router)
|
||||
api_router.include_router(walk_forward_router)
|
||||
api_router.include_router(watchlist_router)
|
||||
api_router.include_router(orders_router)
|
||||
api_router.include_router(real_trades_router)
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
"""API routes for walk-forward backtest optimization (save/load per user).
|
||||
|
||||
See app/services/walk_forward.py for the analysis itself — this module
|
||||
is just the HTTP + persistence layer, mirroring the pattern already used
|
||||
by backtest_history.py for single-run backtests.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from decimal import Decimal
|
||||
from uuid import UUID, uuid4
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import get_db
|
||||
from app.core.deps import get_current_user
|
||||
from app.models.user import User as UserModel
|
||||
from app.services.walk_forward import run_walk_forward
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/walk-forward", tags=["walk_forward"])
|
||||
|
||||
|
||||
@router.post("/run")
|
||||
async def run(
|
||||
symbol: str = Query("BTC/USDT"),
|
||||
exchange: str = Query("mexc"),
|
||||
timeframe: str = Query("4h", description="1h or 4h recommended — lower timeframes multiply the candle count and runtime"),
|
||||
total_days: int = Query(1095, ge=180, le=1825, description="Total lookback in days (default ~3 years)"),
|
||||
train_days: int = Query(270, ge=30, description="Train window size per fold, in days"),
|
||||
test_days: int = Query(90, ge=14, description="Held-out test window size per fold, in days"),
|
||||
trade_size: float = Query(10.0),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_user: UserModel = Depends(get_current_user),
|
||||
):
|
||||
"""Run a walk-forward analysis and persist it to the user's history."""
|
||||
result = await run_walk_forward(
|
||||
db, symbol, exchange, timeframe,
|
||||
total_days=total_days, train_days=train_days, test_days=test_days,
|
||||
trade_size=Decimal(str(trade_size)),
|
||||
)
|
||||
if "error" in result:
|
||||
raise HTTPException(status_code=400, detail=result["error"])
|
||||
|
||||
summary = result["out_of_sample_summary"]
|
||||
wf_id = uuid4()
|
||||
await db.execute(
|
||||
text("""
|
||||
INSERT INTO walk_forward_results
|
||||
(id, user_id, symbol, exchange, timeframe, total_days, train_days, test_days,
|
||||
folds_count, oos_trades, oos_win_rate, oos_total_pnl, oos_profit_factor,
|
||||
oos_max_drawdown_pct, result_json, created_at)
|
||||
VALUES (:id, :uid, :symbol, :exchange, :tf, :total_days, :train_days, :test_days,
|
||||
:folds_count, :oos_trades, :oos_win_rate, :oos_total_pnl, :oos_profit_factor,
|
||||
:oos_max_dd, :rj, :ca)
|
||||
"""),
|
||||
{
|
||||
"id": wf_id, "uid": current_user.id,
|
||||
"symbol": symbol, "exchange": exchange, "tf": timeframe,
|
||||
"total_days": total_days, "train_days": train_days, "test_days": test_days,
|
||||
"folds_count": len(result["folds"]),
|
||||
"oos_trades": summary["trades"], "oos_win_rate": summary["win_rate"],
|
||||
"oos_total_pnl": summary["total_pnl"], "oos_profit_factor": summary["profit_factor"],
|
||||
"oos_max_dd": summary["max_drawdown_pct"],
|
||||
"rj": json.dumps(result, default=str),
|
||||
"ca": datetime.now(timezone.utc),
|
||||
},
|
||||
)
|
||||
await db.commit()
|
||||
|
||||
return {"id": str(wf_id), **result}
|
||||
|
||||
|
||||
@router.get("/history")
|
||||
async def list_walk_forward_runs(
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_user: UserModel = Depends(get_current_user),
|
||||
):
|
||||
"""List walk-forward run summaries for the current user (no fold detail)."""
|
||||
result = await db.execute(
|
||||
text("""
|
||||
SELECT id, symbol, exchange, timeframe, total_days, train_days, test_days,
|
||||
folds_count, oos_trades, oos_win_rate, oos_total_pnl,
|
||||
oos_profit_factor, oos_max_drawdown_pct, created_at
|
||||
FROM walk_forward_results
|
||||
WHERE user_id = :uid
|
||||
ORDER BY created_at DESC
|
||||
LIMIT :limit
|
||||
"""),
|
||||
{"uid": current_user.id, "limit": limit},
|
||||
)
|
||||
rows = result.fetchall()
|
||||
return [
|
||||
{
|
||||
"id": str(r[0]),
|
||||
"symbol": r[1], "exchange": r[2], "timeframe": r[3],
|
||||
"total_days": r[4], "train_days": r[5], "test_days": r[6],
|
||||
"folds_count": r[7], "oos_trades": r[8],
|
||||
"oos_win_rate": float(r[9]) if r[9] is not None else None,
|
||||
"oos_total_pnl": float(r[10]) if r[10] is not None else None,
|
||||
"oos_profit_factor": float(r[11]) if r[11] is not None else None,
|
||||
"oos_max_drawdown_pct": float(r[12]) if r[12] is not None else None,
|
||||
"created_at": r[13].isoformat() if r[13] else None,
|
||||
}
|
||||
for r in rows
|
||||
]
|
||||
|
||||
|
||||
@router.get("/{wf_id}")
|
||||
async def get_walk_forward_run(
|
||||
wf_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_user: UserModel = Depends(get_current_user),
|
||||
):
|
||||
"""Return the full fold-by-fold detail for one saved walk-forward run."""
|
||||
try:
|
||||
wf_uuid = UUID(wf_id)
|
||||
except ValueError:
|
||||
raise HTTPException(400, "Invalid ID")
|
||||
|
||||
result = await db.execute(
|
||||
text("SELECT id, result_json, created_at FROM walk_forward_results WHERE id = :id AND user_id = :uid"),
|
||||
{"id": wf_uuid, "uid": current_user.id},
|
||||
)
|
||||
row = result.fetchone()
|
||||
if not row:
|
||||
raise HTTPException(404, "Walk-forward run not found")
|
||||
|
||||
detail = json.loads(row[1])
|
||||
detail["id"] = str(row[0])
|
||||
detail["created_at"] = row[2].isoformat() if row[2] else None
|
||||
return detail
|
||||
|
||||
|
||||
@router.delete("/{wf_id}")
|
||||
async def delete_walk_forward_run(
|
||||
wf_id: str,
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_user: UserModel = Depends(get_current_user),
|
||||
):
|
||||
"""Delete a saved walk-forward run."""
|
||||
try:
|
||||
wf_uuid = UUID(wf_id)
|
||||
except ValueError:
|
||||
raise HTTPException(400, "Invalid ID")
|
||||
|
||||
result = await db.execute(
|
||||
text("DELETE FROM walk_forward_results WHERE id = :id AND user_id = :uid"),
|
||||
{"id": wf_uuid, "uid": current_user.id},
|
||||
)
|
||||
await db.commit()
|
||||
if result.rowcount == 0:
|
||||
raise HTTPException(404, "Walk-forward run not found")
|
||||
return {"message": "Deleted"}
|
||||
@@ -0,0 +1,416 @@
|
||||
"""Backtest engine — candle fetch, indicator precompute, and trade simulation.
|
||||
|
||||
Split out of `app/api/v1/backtest.py` (which now only holds the FastAPI
|
||||
routes) so both the single-run `/backtest/run` endpoint and the
|
||||
walk-forward optimizer (`app/services/walk_forward.py`) share one
|
||||
implementation instead of two copies drifting apart.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import select, and_
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.candle import Candle
|
||||
from app.models.symbol import Symbol
|
||||
from app.models.exchange import Exchange
|
||||
from app.services.indicator_service import (
|
||||
bollinger_bands, rsi, sma, macd, supertrend,
|
||||
volume_breakout, ichimoku, detect_divergence, market_structure,
|
||||
)
|
||||
from app.services.signal_scoring import (
|
||||
_classify_signal_combined,
|
||||
STRONG_BUY, BUY, STRONG_SELL, SELL,
|
||||
)
|
||||
|
||||
# Min candles for warmup: BB(20) + RSI(14) + some room = 30
|
||||
MIN_CANDLES = 30
|
||||
|
||||
|
||||
async def _fetch_symbol(db: AsyncSession, symbol: str, exchange: str) -> Symbol | None:
|
||||
"""Look up a Symbol row by (symbol, exchange) name."""
|
||||
result = await db.execute(
|
||||
select(Symbol)
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(and_(Exchange.name == exchange, Symbol.symbol == symbol))
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def _fetch_candles(
|
||||
db: AsyncSession,
|
||||
symbol_id,
|
||||
timeframe: str,
|
||||
since: datetime,
|
||||
until: datetime | None = None,
|
||||
) -> list[Candle]:
|
||||
"""Fetch candles for a symbol/timeframe in [since, until), ordered ascending."""
|
||||
conditions = [
|
||||
Candle.symbol_id == symbol_id,
|
||||
Candle.timeframe == timeframe,
|
||||
Candle.timestamp >= since,
|
||||
]
|
||||
if until is not None:
|
||||
conditions.append(Candle.timestamp < until)
|
||||
result = await db.execute(
|
||||
select(Candle).where(and_(*conditions)).order_by(Candle.timestamp.asc())
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
def _precompute_indicators(candles: list[Candle], timeframe: str) -> dict:
|
||||
"""Pre-compute all candle data & indicators ONCE for a candle range.
|
||||
|
||||
Split out so `walk_forward.py` can precompute indicators for a fold's
|
||||
data once, then cheaply replay many threshold combinations against it
|
||||
(see `_simulate_trades`).
|
||||
"""
|
||||
# MTF config
|
||||
tf_minutes = {"15m": 15, "30m": 30, "1h": 60, "4h": 240}
|
||||
main_minutes = tf_minutes.get(timeframe, 30)
|
||||
mtf_config = []
|
||||
for mtf_tf, mtf_minutes, mtf_w in [("15m", 15, 0.5), ("1h", 60, 1.5), ("4h", 240, 2.0)]:
|
||||
if mtf_tf == timeframe:
|
||||
continue
|
||||
mult = mtf_minutes // main_minutes
|
||||
if mult >= 1 and len(candles) >= mult * MIN_CANDLES:
|
||||
mtf_config.append((mtf_tf, mult, mtf_w))
|
||||
|
||||
candle_dicts_full = [
|
||||
{"high": float(c.high), "low": float(c.low),
|
||||
"close": float(c.close), "open": float(c.open),
|
||||
"volume": float(c.volume)}
|
||||
for c in candles
|
||||
]
|
||||
close_prices_full = [float(c.close) for c in candles]
|
||||
|
||||
# Pre-compute indicators on full dataset (O(n) instead of O(n²))
|
||||
bb_full = bollinger_bands(close_prices_full) or {}
|
||||
rsi_full = rsi(close_prices_full) or []
|
||||
sma_full = sma(close_prices_full, 20) or []
|
||||
macd_full = macd(close_prices_full) or {}
|
||||
st_full = supertrend(candle_dicts_full) or {}
|
||||
vb_full = volume_breakout(candle_dicts_full) or []
|
||||
ichi_full = ichimoku(candle_dicts_full) or {}
|
||||
smc_full = market_structure(candle_dicts_full) or {}
|
||||
|
||||
# Pre-compute divergence ONCE (uses full arrays, indexes match)
|
||||
rsi_div_full = detect_divergence(close_prices_full, rsi_full)
|
||||
macd_hist_full = macd_full.get("histogram", []) if macd_full else []
|
||||
macd_div_full = detect_divergence(close_prices_full, macd_hist_full)
|
||||
|
||||
# Pre-build MTF candles ONCE per MTF config
|
||||
mtf_precomputed = []
|
||||
for mtf_name, mtf_mult, mtf_w in mtf_config:
|
||||
mtf_candles_list = []
|
||||
for j in range(0, len(candle_dicts_full) - mtf_mult + 1, mtf_mult):
|
||||
chunk = candle_dicts_full[j:j + mtf_mult]
|
||||
mtf_candles_list.append({
|
||||
"open": chunk[0]["open"],
|
||||
"high": max(c["high"] for c in chunk),
|
||||
"low": min(c["low"] for c in chunk),
|
||||
"close": chunk[-1]["close"],
|
||||
"volume": sum(c["volume"] for c in chunk),
|
||||
})
|
||||
if len(mtf_candles_list) >= MIN_CANDLES:
|
||||
mtf_p = [c["close"] for c in mtf_candles_list]
|
||||
mtf_precomputed.append({
|
||||
"name": mtf_name,
|
||||
"weight": mtf_w,
|
||||
"mult": mtf_mult,
|
||||
"candles_list": mtf_candles_list,
|
||||
"close_prices": mtf_p,
|
||||
"bb": bollinger_bands(mtf_p) or {},
|
||||
"rsi": rsi(mtf_p) or [],
|
||||
"sma": sma(mtf_p, 20) or [],
|
||||
"macd": macd(mtf_p) or {},
|
||||
"st": supertrend(mtf_candles_list) or {},
|
||||
"vb": volume_breakout(mtf_candles_list) or [],
|
||||
"ichi": ichimoku(mtf_candles_list) or {},
|
||||
"smc": market_structure(mtf_candles_list) or {},
|
||||
})
|
||||
|
||||
return {
|
||||
"candle_dicts_full": candle_dicts_full,
|
||||
"close_prices_full": close_prices_full,
|
||||
"bb_full": bb_full,
|
||||
"rsi_full": rsi_full,
|
||||
"sma_full": sma_full,
|
||||
"macd_full": macd_full,
|
||||
"st_full": st_full,
|
||||
"vb_full": vb_full,
|
||||
"ichi_full": ichi_full,
|
||||
"smc_full": smc_full,
|
||||
"rsi_div_full": rsi_div_full,
|
||||
"macd_div_full": macd_div_full,
|
||||
"mtf_precomputed": mtf_precomputed,
|
||||
}
|
||||
|
||||
|
||||
def _simulate_trades(
|
||||
candles: list[Candle],
|
||||
precomputed: dict,
|
||||
trade_size: Decimal = Decimal("10"),
|
||||
strong_threshold: float = 4.0,
|
||||
signal_threshold: float = 1.0,
|
||||
max_hold_candles: int = 48,
|
||||
active_from_index: int = MIN_CANDLES,
|
||||
) -> tuple[list[dict], list[dict]]:
|
||||
"""Replay signal classification + trade simulation over precomputed indicators.
|
||||
|
||||
`active_from_index` lets callers pass extra warmup candles before the
|
||||
window they actually want simulated (e.g. walk-forward fold
|
||||
boundaries) — candles before this index are used only so indicators
|
||||
have enough lookback, never turned into signals/trades.
|
||||
"""
|
||||
close_prices_full = precomputed["close_prices_full"]
|
||||
bb_full = precomputed["bb_full"]
|
||||
rsi_full = precomputed["rsi_full"]
|
||||
sma_full = precomputed["sma_full"]
|
||||
macd_full = precomputed["macd_full"]
|
||||
st_full = precomputed["st_full"]
|
||||
vb_full = precomputed["vb_full"]
|
||||
ichi_full = precomputed["ichi_full"]
|
||||
smc_full = precomputed["smc_full"]
|
||||
rsi_div_full = precomputed["rsi_div_full"]
|
||||
macd_div_full = precomputed["macd_div_full"]
|
||||
mtf_precomputed = precomputed["mtf_precomputed"]
|
||||
|
||||
all_signals = []
|
||||
trades = []
|
||||
current_position = None
|
||||
|
||||
start_index = max(MIN_CANDLES, active_from_index)
|
||||
for i in range(start_index, len(candles)):
|
||||
candle = candles[i]
|
||||
latest_close = close_prices_full[i]
|
||||
timestamp = candle.timestamp.isoformat()
|
||||
|
||||
# Slice pre-computed arrays (O(i) but ~100x faster than recomputing)
|
||||
clip = i + 1
|
||||
def _safe_slice(v):
|
||||
return v[:clip] if v is not None and hasattr(v, '__getitem__') else v
|
||||
bb_data = {k: _safe_slice(v) for k, v in bb_full.items()} if bb_full else {}
|
||||
rsi_data = rsi_full[:clip] if rsi_full else []
|
||||
sma_data = sma_full[:clip] if sma_full else []
|
||||
macd_data = {k: _safe_slice(v) for k, v in macd_full.items()} if macd_full else {}
|
||||
st_data = {k: _safe_slice(v) for k, v in st_full.items()} if st_full else {}
|
||||
vb_data = vb_full[:clip] if vb_full else []
|
||||
ichi_data = {k: _safe_slice(v) for k, v in ichi_full.items()} if ichi_full else {}
|
||||
smc_data = {k: _safe_slice(v) for k, v in smc_full.items()} if smc_full else {}
|
||||
|
||||
# MTF votes — use precomputed MTF indicators, sliced to current MTF candle index
|
||||
mtf_votes = []
|
||||
for mtf in mtf_precomputed:
|
||||
# Which MTF candle corresponds to main candle i?
|
||||
mtf_idx = i // mtf["mult"]
|
||||
if mtf_idx < MIN_CANDLES or mtf_idx >= len(mtf["close_prices"]):
|
||||
continue
|
||||
clip_mtf = mtf_idx + 1
|
||||
def _safe_slice_mtf(v):
|
||||
return v[:clip_mtf] if v is not None and hasattr(v, '__getitem__') else v
|
||||
mtf_s, *_ = _classify_signal_combined(
|
||||
mtf["close_prices"][mtf_idx],
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["bb"].items()},
|
||||
mtf["rsi"][:clip_mtf],
|
||||
mtf["sma"][:clip_mtf],
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["macd"].items()} if mtf["macd"] else None,
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["st"].items()} if mtf["st"] else None,
|
||||
mtf["vb"][:clip_mtf] if mtf["vb"] else None,
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["ichi"].items()} if mtf["ichi"] else None,
|
||||
(None, None), (None, None),
|
||||
{k: _safe_slice_mtf(v) for k, v in mtf["smc"].items()} if mtf["smc"] else None,
|
||||
)
|
||||
if mtf_s:
|
||||
mtf_votes.append((mtf_s, "", mtf["weight"]))
|
||||
|
||||
signal_type, strength, *_ = _classify_signal_combined(
|
||||
latest_close, bb_data, rsi_data, sma_data,
|
||||
macd_data, st_data, vb_data, ichi_data,
|
||||
rsi_div_full, macd_div_full, smc_data, mtf_votes or None,
|
||||
strong_threshold=strong_threshold, signal_threshold=signal_threshold,
|
||||
)
|
||||
|
||||
if signal_type:
|
||||
all_signals.append({
|
||||
"time": timestamp, "signal": signal_type,
|
||||
"strength": strength or "", "price": latest_close,
|
||||
})
|
||||
|
||||
# PnL simulation
|
||||
if signal_type in (STRONG_BUY, BUY):
|
||||
if current_position and current_position["direction"] == "SHORT":
|
||||
if signal_type == STRONG_BUY:
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
pnl = (entry - latest_close) * qty
|
||||
current_position.update({
|
||||
"exit_price": latest_close, "exit_time": timestamp,
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
|
||||
})
|
||||
trades.append(current_position)
|
||||
current_position = None
|
||||
else:
|
||||
continue
|
||||
if not current_position:
|
||||
qty = float(trade_size) / latest_close
|
||||
current_position = {
|
||||
"direction": "LONG", "entry_price": latest_close,
|
||||
"entry_time": timestamp, "quantity": qty,
|
||||
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
|
||||
}
|
||||
|
||||
elif signal_type in (STRONG_SELL, SELL):
|
||||
if current_position and current_position["direction"] == "LONG":
|
||||
if signal_type == STRONG_SELL:
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
pnl = (latest_close - entry) * qty
|
||||
current_position.update({
|
||||
"exit_price": latest_close, "exit_time": timestamp,
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
|
||||
})
|
||||
trades.append(current_position)
|
||||
current_position = None
|
||||
else:
|
||||
continue
|
||||
if not current_position:
|
||||
qty = float(trade_size) / latest_close
|
||||
current_position = {
|
||||
"direction": "SHORT", "entry_price": latest_close,
|
||||
"entry_time": timestamp, "quantity": qty,
|
||||
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
|
||||
}
|
||||
|
||||
# Time limit
|
||||
if current_position and current_position["status"] == "OPEN":
|
||||
hold = i - current_position["entry_index"]
|
||||
if hold >= max_hold_candles:
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
if current_position["direction"] == "LONG":
|
||||
pnl = (latest_close - entry) * qty
|
||||
else:
|
||||
pnl = (entry - latest_close) * qty
|
||||
current_position.update({
|
||||
"exit_price": latest_close, "exit_time": timestamp,
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "TIME_LIMIT",
|
||||
})
|
||||
trades.append(current_position)
|
||||
current_position = None
|
||||
|
||||
# Close final position
|
||||
if current_position and current_position["status"] == "OPEN":
|
||||
last_close = float(candles[-1].close)
|
||||
entry = current_position["entry_price"]
|
||||
qty = current_position["quantity"]
|
||||
if current_position["direction"] == "LONG":
|
||||
pnl = (last_close - entry) * qty
|
||||
else:
|
||||
pnl = (entry - last_close) * qty
|
||||
current_position.update({
|
||||
"exit_price": last_close,
|
||||
"exit_time": candles[-1].timestamp.isoformat(),
|
||||
"pnl": pnl, "status": "CLOSED", "exit_reason": "END_OF_DATA",
|
||||
})
|
||||
trades.append(current_position)
|
||||
|
||||
return all_signals, trades
|
||||
|
||||
|
||||
def _compute_stats(all_signals: list[dict], trades: list[dict]) -> dict:
|
||||
"""Reduce raw signals/trades into the summary stats block used by
|
||||
both the single-run backtest and each walk-forward fold."""
|
||||
counts = defaultdict(int)
|
||||
for s in all_signals:
|
||||
counts[s["signal"]] += 1
|
||||
|
||||
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
|
||||
winning_trades = [t for t in closed_trades if t.get("pnl", 0) > 0]
|
||||
losing_trades = [t for t in closed_trades if t.get("pnl", 0) <= 0]
|
||||
total_pnl = sum(t.get("pnl", 0) for t in closed_trades)
|
||||
gross_profit = sum(t.get("pnl", 0) for t in winning_trades)
|
||||
gross_loss = sum(t.get("pnl", 0) for t in losing_trades)
|
||||
win_rate = round(len(winning_trades) / len(closed_trades) * 100, 1) if closed_trades else 0
|
||||
profit_factor = round(abs(gross_profit / gross_loss), 2) if gross_loss != 0 else None
|
||||
avg_win = round(gross_profit / len(winning_trades), 2) if winning_trades else None
|
||||
avg_loss = round(gross_loss / len(losing_trades), 2) if losing_trades else None
|
||||
|
||||
best_trade = max(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
|
||||
worst_trade = min(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
|
||||
|
||||
return {
|
||||
"signal_counts": dict(counts),
|
||||
"total_signals": len(all_signals),
|
||||
"recent_signals": all_signals[-15:],
|
||||
"trades": {
|
||||
"total": len(closed_trades),
|
||||
"wins": len(winning_trades),
|
||||
"losses": len(losing_trades),
|
||||
"win_rate": win_rate,
|
||||
"total_pnl": round(total_pnl, 2),
|
||||
"profit_factor": profit_factor,
|
||||
"avg_win": avg_win,
|
||||
"avg_loss": avg_loss,
|
||||
"best_trade": {
|
||||
"direction": best_trade.get("direction"),
|
||||
"entry_price": round(best_trade["entry_price"], 4),
|
||||
"exit_price": round(best_trade["exit_price"], 4),
|
||||
"pnl": round(best_trade["pnl"], 2),
|
||||
"entry_signal": best_trade.get("entry_signal"),
|
||||
} if best_trade else None,
|
||||
"worst_trade": {
|
||||
"direction": worst_trade.get("direction"),
|
||||
"entry_price": round(worst_trade["entry_price"], 4),
|
||||
"exit_price": round(worst_trade["exit_price"], 4),
|
||||
"pnl": round(worst_trade["pnl"], 2),
|
||||
"entry_signal": worst_trade.get("entry_signal"),
|
||||
} if worst_trade else None,
|
||||
"per_signal": {},
|
||||
"recent": closed_trades[-10:],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
async def run_backtest(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
exchange: str,
|
||||
timeframe: str = "30m",
|
||||
days: int = 7,
|
||||
trade_size: Decimal = Decimal("10"),
|
||||
) -> dict:
|
||||
"""Run a single backtest over the last `days` days and return structured results."""
|
||||
db_symbol = await _fetch_symbol(db, symbol, exchange)
|
||||
if not db_symbol:
|
||||
return {"error": f"Symbol {symbol} not found on {exchange}"}
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
|
||||
candles = await _fetch_candles(db, db_symbol.id, timeframe, since=cutoff)
|
||||
|
||||
if len(candles) < MIN_CANDLES:
|
||||
if candles:
|
||||
span_hours = (candles[-1].timestamp - candles[0].timestamp).total_seconds() / 3600
|
||||
if span_hours >= 24:
|
||||
avail = f"{span_hours/24:.0f}d"
|
||||
else:
|
||||
avail = f"{span_hours:.0f}h"
|
||||
return {"error": f"Only ~{avail} data available (need at least {MIN_CANDLES} candles for {timeframe}). Try more days or a higher timeframe (4h)."}
|
||||
return {"error": f"No candle data found for {symbol} on {timeframe}. The exchange may not support this pair."}
|
||||
|
||||
precomputed = _precompute_indicators(candles, timeframe)
|
||||
all_signals, trades = _simulate_trades(candles, precomputed, trade_size)
|
||||
stats = _compute_stats(all_signals, trades)
|
||||
|
||||
return {
|
||||
"symbol": symbol,
|
||||
"exchange": exchange,
|
||||
"timeframe": timeframe,
|
||||
"days": days,
|
||||
"candles_count": len(candles),
|
||||
**stats,
|
||||
}
|
||||
@@ -165,7 +165,7 @@ def _classify_signal_bb(
|
||||
return None, None
|
||||
|
||||
|
||||
def _classify_signal_combined(
|
||||
def _compute_adjusted_score(
|
||||
close_price: float,
|
||||
bb: dict[str, list[float]],
|
||||
rsi: list[float] | None,
|
||||
@@ -185,47 +185,37 @@ def _classify_signal_combined(
|
||||
candlestick_score: float | None = None,
|
||||
rates: dict[str, float] | None = None,
|
||||
enabled_strategies: list[str] | None = None,
|
||||
) -> tuple[Optional[str], Optional[str], float, dict[str, float]]:
|
||||
"""Classify market state using 13-algorithm voting with win-rate boosting.
|
||||
) -> tuple[Optional[str], Optional[str], float, float, dict[str, float]]:
|
||||
"""Run the 13-algorithm vote and reduce it to a single adjusted score.
|
||||
|
||||
Algorithms:
|
||||
1. Double BB + RSI
|
||||
2. MACD Crossover
|
||||
3. SuperTrend
|
||||
4. Volume Breakout
|
||||
5. Ichimoku Cloud
|
||||
6. Divergence Detection (RSI + MACD)
|
||||
7. 🌤️ Market Structure (SMC) — BOS, CHoCH, OB
|
||||
8. 🔄 Multi-Timeframe (15m + 1h + 4h)
|
||||
9. 📊 OBV (On-Balance Volume) Crossover
|
||||
10. 🔄 Stochastic RSI Crossover
|
||||
11. 💰 MFI (Money Flow Index)
|
||||
12. 🕯️ FVG (Fair Value Gap)
|
||||
13. 🕯️ Candlestick Patterns (30+ patterns)
|
||||
This is the expensive, threshold-independent half of signal
|
||||
classification — algorithms 1-13, win-rate boosting, correlation
|
||||
dampening, and dynamic normalization. It does NOT decide the final
|
||||
signal type; that is a cheap final step in `_classify_signal_combined`
|
||||
(or `_score_to_signal`) so callers that need to try many threshold
|
||||
combinations (e.g. walk-forward parameter search) can compute this
|
||||
once per candle and replay different thresholds against it cheaply.
|
||||
|
||||
Each algorithm votes: BUY (+1/+2), SELL (-1/-2), or NEUTRAL (0).
|
||||
If *rates* is provided, each strategy's raw score is boosted by its
|
||||
historical win rate before the final classification.
|
||||
|
||||
Returns (signal_type, strength, confidence, raw_scores) where
|
||||
confidence is a 0-1 float and raw_scores is a dict of all 9
|
||||
algorithm scores for ML feature collection.
|
||||
Returns (override_signal, override_strength, adjusted_score, confidence,
|
||||
raw_scores). When override_signal is not None (SQUEEZE_ALERT,
|
||||
CAUTION_LONG, CAUTION_SHORT), the caller must return it as-is —
|
||||
it bypasses threshold-based classification entirely.
|
||||
"""
|
||||
# ── NaN/Inf guard: reject any invalid price before processing ──
|
||||
if not math.isfinite(close_price) or close_price <= 0:
|
||||
logger.warning("_classify_signal_combined: invalid close_price=%s, returning NEUTRAL", close_price)
|
||||
return None, None, 0.0, {}
|
||||
return None, None, 0.0, 0.0, {}
|
||||
|
||||
# ── Special signals (override) ──
|
||||
squeeze = _detect_squeeze(bb)
|
||||
if squeeze:
|
||||
return SQUEEZE_ALERT, "MODERATE", 0.5, {}
|
||||
return SQUEEZE_ALERT, "MODERATE", 0.0, 0.5, {}
|
||||
|
||||
# P2-2: Call _classify_signal_bb ONCE, reuse result for both
|
||||
# early-return check AND the raw_scores vote
|
||||
bb_type, bb_strength = _classify_signal_bb(close_price, bb, rsi, sma)
|
||||
if bb_type in (CAUTION_LONG, CAUTION_SHORT):
|
||||
return bb_type, "MODERATE", 0.5, {}
|
||||
return bb_type, "MODERATE", 0.0, 0.5, {}
|
||||
|
||||
# ── Collect per-strategy raw scores ──
|
||||
raw_scores: dict[str, float] = {
|
||||
@@ -504,18 +494,95 @@ def _classify_signal_combined(
|
||||
else:
|
||||
adjusted_score = total_score
|
||||
|
||||
# ── Final classification from boosted score ──
|
||||
# 🔧 Dynamic thresholds: STRONG needs effective 4.0, BUY/SELL needs 1.0
|
||||
if adjusted_score >= 4.0:
|
||||
return STRONG_BUY, "STRONG", confidence, raw_scores
|
||||
elif adjusted_score >= 1.0:
|
||||
return BUY, "MODERATE", confidence, raw_scores
|
||||
elif adjusted_score <= -4.0:
|
||||
return STRONG_SELL, "STRONG", confidence, raw_scores
|
||||
elif adjusted_score <= -1.0:
|
||||
return SELL, "MODERATE", confidence, raw_scores
|
||||
return None, None, adjusted_score, confidence, raw_scores
|
||||
|
||||
return None, None, confidence, raw_scores
|
||||
|
||||
def _score_to_signal(
|
||||
adjusted_score: float,
|
||||
strong_threshold: float = 4.0,
|
||||
signal_threshold: float = 1.0,
|
||||
) -> tuple[Optional[str], Optional[str]]:
|
||||
"""Turn an adjusted score into a signal type — the cheap, threshold-only
|
||||
half of classification. Split out from `_compute_adjusted_score` so
|
||||
walk-forward parameter search can replay many threshold combinations
|
||||
against an already-computed score array without re-running the 13
|
||||
algorithms each time.
|
||||
|
||||
🔧 Dynamic thresholds: STRONG needs effective `strong_threshold` (default
|
||||
4.0), BUY/SELL needs `signal_threshold` (default 1.0).
|
||||
"""
|
||||
if adjusted_score >= strong_threshold:
|
||||
return STRONG_BUY, "STRONG"
|
||||
elif adjusted_score >= signal_threshold:
|
||||
return BUY, "MODERATE"
|
||||
elif adjusted_score <= -strong_threshold:
|
||||
return STRONG_SELL, "STRONG"
|
||||
elif adjusted_score <= -signal_threshold:
|
||||
return SELL, "MODERATE"
|
||||
return None, None
|
||||
|
||||
|
||||
def _classify_signal_combined(
|
||||
close_price: float,
|
||||
bb: dict[str, list[float]],
|
||||
rsi: list[float] | None,
|
||||
sma: list[float] | None,
|
||||
macd_data: dict | None,
|
||||
st_data: dict | None,
|
||||
vol_data: list | None,
|
||||
ichi_data: dict | None = None,
|
||||
rsi_div: tuple = (None, None),
|
||||
macd_div: tuple = (None, None),
|
||||
smc_data: dict | None = None,
|
||||
mtf_votes: list[tuple[Optional[str], Optional[str], float]] | None = None,
|
||||
obv_data: list | None = None,
|
||||
stoch_rsi_data: dict | None = None,
|
||||
mfi_data: list | None = None,
|
||||
fvg_data: dict | None = None,
|
||||
candlestick_score: float | None = None,
|
||||
rates: dict[str, float] | None = None,
|
||||
enabled_strategies: list[str] | None = None,
|
||||
strong_threshold: float = 4.0,
|
||||
signal_threshold: float = 1.0,
|
||||
) -> tuple[Optional[str], Optional[str], float, dict[str, float]]:
|
||||
"""Classify market state using 13-algorithm voting with win-rate boosting.
|
||||
|
||||
Algorithms:
|
||||
1. Double BB + RSI
|
||||
2. MACD Crossover
|
||||
3. SuperTrend
|
||||
4. Volume Breakout
|
||||
5. Ichimoku Cloud
|
||||
6. Divergence Detection (RSI + MACD)
|
||||
7. 🌤️ Market Structure (SMC) — BOS, CHoCH, OB
|
||||
8. 🔄 Multi-Timeframe (15m + 1h + 4h)
|
||||
9. 📊 OBV (On-Balance Volume) Crossover
|
||||
10. 🔄 Stochastic RSI Crossover
|
||||
11. 💰 MFI (Money Flow Index)
|
||||
12. 🕯️ FVG (Fair Value Gap)
|
||||
13. 🕯️ Candlestick Patterns (30+ patterns)
|
||||
|
||||
Each algorithm votes: BUY (+1/+2), SELL (-1/-2), or NEUTRAL (0).
|
||||
If *rates* is provided, each strategy's raw score is boosted by its
|
||||
historical win rate before the final classification. *strong_threshold*
|
||||
and *signal_threshold* control the final cutoffs (see `_score_to_signal`)
|
||||
— left at their defaults for live trading; walk-forward backtesting
|
||||
overrides them during parameter search.
|
||||
|
||||
Returns (signal_type, strength, confidence, raw_scores) where
|
||||
confidence is a 0-1 float and raw_scores is a dict of all 9
|
||||
algorithm scores for ML feature collection.
|
||||
"""
|
||||
override_signal, override_strength, adjusted_score, confidence, raw_scores = _compute_adjusted_score(
|
||||
close_price, bb, rsi, sma, macd_data, st_data, vol_data, ichi_data,
|
||||
rsi_div, macd_div, smc_data, mtf_votes, obv_data, stoch_rsi_data,
|
||||
mfi_data, fvg_data, candlestick_score, rates, enabled_strategies,
|
||||
)
|
||||
if override_signal is not None:
|
||||
return override_signal, override_strength, confidence, raw_scores
|
||||
|
||||
signal_type, strength = _score_to_signal(adjusted_score, strong_threshold, signal_threshold)
|
||||
return signal_type, strength, confidence, raw_scores
|
||||
|
||||
|
||||
def _calculate_pnl(
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
"""Walk-forward backtest optimization.
|
||||
|
||||
Splits historical data into rolling train/test folds, auto-optimizes a
|
||||
small parameter grid on each fold's train window, then evaluates the
|
||||
optimized parameters on that fold's held-out test window
|
||||
(out-of-sample). Stitching all out-of-sample test results together
|
||||
gives an honest performance estimate that isn't inflated by tuning
|
||||
parameters against the same data used to score them — see item (m) in
|
||||
theo_doi_trading-portal_v6.md.
|
||||
|
||||
Only the three cheaply-tunable "when to enter/exit" parameters are
|
||||
optimized (see `_score_to_signal` in signal_scoring.py): the two score
|
||||
thresholds and the max hold time. The 13-algorithm voting internals
|
||||
(RSI/MFI/etc. cutoffs) are not parameterized — doing so would require a
|
||||
much larger, riskier refactor of signal_scoring.py.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from itertools import product
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.services.backtest_engine import (
|
||||
MIN_CANDLES,
|
||||
_fetch_symbol,
|
||||
_fetch_candles,
|
||||
_precompute_indicators,
|
||||
_simulate_trades,
|
||||
_compute_stats,
|
||||
)
|
||||
|
||||
# Parameters that are cheap to grid-search: they only affect the final
|
||||
# threshold/exit logic, not the 13-algorithm scoring itself, so replaying
|
||||
# them against already-precomputed indicators is fast (see _simulate_trades).
|
||||
DEFAULT_PARAM_GRID: dict[str, list[float]] = {
|
||||
"strong_threshold": [3.5, 4.0, 4.5],
|
||||
"signal_threshold": [0.75, 1.0, 1.5],
|
||||
"max_hold_candles": [24, 48, 96],
|
||||
}
|
||||
|
||||
MIN_TRADES_PER_FOLD = 5 # reject param combos too sparse to trust
|
||||
WARMUP_BUFFER_CANDLES = 60 # extra history fetched before each window so indicators aren't cold at window start
|
||||
|
||||
_TF_MINUTES = {"15m": 15, "30m": 30, "1h": 60, "4h": 240, "1d": 1440}
|
||||
|
||||
|
||||
def _fold_score(closed_trades: list[dict], min_trades: int = MIN_TRADES_PER_FOLD) -> float:
|
||||
"""Per-trade Sharpe-like objective, scaled by sqrt(n) trades.
|
||||
|
||||
Favors consistent edge over one lucky trade, and rejects parameter
|
||||
combinations with too few trades to be statistically meaningful
|
||||
(returns -inf so they never win a grid search).
|
||||
"""
|
||||
if len(closed_trades) < min_trades:
|
||||
return float("-inf")
|
||||
pnls = [float(t.get("pnl", 0.0)) for t in closed_trades]
|
||||
n = len(pnls)
|
||||
mean = sum(pnls) / n
|
||||
variance = sum((p - mean) ** 2 for p in pnls) / n
|
||||
std = math.sqrt(variance)
|
||||
if std == 0:
|
||||
return mean * math.sqrt(n)
|
||||
return (mean / std) * math.sqrt(n)
|
||||
|
||||
|
||||
def generate_folds(
|
||||
total_days: int,
|
||||
train_days: int,
|
||||
test_days: int,
|
||||
anchor: datetime | None = None,
|
||||
) -> list[dict]:
|
||||
"""Generate rolling folds: a fixed-size train window sliding forward
|
||||
by `test_days` each step, anchored to `anchor` (default: now) counting
|
||||
back `total_days`. Each fold covers [train_start, train_end) train +
|
||||
[train_end, test_end) test, walked forward in chronological order.
|
||||
"""
|
||||
anchor = anchor or datetime.now(timezone.utc)
|
||||
origin = anchor - timedelta(days=total_days)
|
||||
|
||||
folds = []
|
||||
offset = 0
|
||||
while True:
|
||||
train_start = origin + timedelta(days=offset)
|
||||
train_end = train_start + timedelta(days=train_days)
|
||||
test_end = train_end + timedelta(days=test_days)
|
||||
if test_end > anchor:
|
||||
break
|
||||
folds.append({
|
||||
"fold_index": len(folds),
|
||||
"train_start": train_start,
|
||||
"train_end": train_end,
|
||||
"test_start": train_end,
|
||||
"test_end": test_end,
|
||||
})
|
||||
offset += test_days
|
||||
return folds
|
||||
|
||||
|
||||
def _warmup_days(timeframe: str, buffer_candles: int = WARMUP_BUFFER_CANDLES) -> int:
|
||||
minutes = _TF_MINUTES.get(timeframe, 30)
|
||||
return max(1, math.ceil(buffer_candles * minutes / 1440))
|
||||
|
||||
|
||||
async def _prepare_window(
|
||||
db: AsyncSession,
|
||||
symbol_id,
|
||||
timeframe: str,
|
||||
window_start: datetime,
|
||||
window_end: datetime,
|
||||
) -> tuple[list, dict, int] | None:
|
||||
"""Fetch candles for [window_start - warmup, window_end) and precompute
|
||||
indicators. Returns (candles, precomputed, active_from_index), where
|
||||
active_from_index is the candle index at which window_start begins —
|
||||
candles before it exist only to warm up indicators, and are never
|
||||
turned into signals/trades. Returns None if there isn't enough data.
|
||||
"""
|
||||
warmup_start = window_start - timedelta(days=_warmup_days(timeframe))
|
||||
candles = await _fetch_candles(db, symbol_id, timeframe, since=warmup_start, until=window_end)
|
||||
if len(candles) < MIN_CANDLES:
|
||||
return None
|
||||
|
||||
active_from_index = None
|
||||
for idx, c in enumerate(candles):
|
||||
if c.timestamp >= window_start:
|
||||
active_from_index = idx
|
||||
break
|
||||
if active_from_index is None:
|
||||
return None # no candles actually within the window itself
|
||||
|
||||
precomputed = _precompute_indicators(candles, timeframe)
|
||||
return candles, precomputed, active_from_index
|
||||
|
||||
|
||||
def _run_combo(candles, precomputed, trade_size, params, active_from_index):
|
||||
return _simulate_trades(
|
||||
candles, precomputed, trade_size,
|
||||
strong_threshold=params["strong_threshold"],
|
||||
signal_threshold=params["signal_threshold"],
|
||||
max_hold_candles=int(params["max_hold_candles"]),
|
||||
active_from_index=active_from_index,
|
||||
)
|
||||
|
||||
|
||||
def _grid_search(
|
||||
candles: list,
|
||||
precomputed: dict,
|
||||
param_grid: dict[str, list[float]],
|
||||
trade_size: Decimal,
|
||||
active_from_index: int,
|
||||
) -> tuple[dict[str, float], float, dict]:
|
||||
"""Try every combination in param_grid, return the one that scores best
|
||||
on the train window by `_fold_score`."""
|
||||
keys = list(param_grid.keys())
|
||||
best_params: dict[str, float] | None = None
|
||||
best_score = float("-inf")
|
||||
best_stats: dict = {}
|
||||
|
||||
for combo in product(*(param_grid[k] for k in keys)):
|
||||
params = dict(zip(keys, combo))
|
||||
all_signals, trades = _run_combo(candles, precomputed, trade_size, params, active_from_index)
|
||||
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
|
||||
score = _fold_score(closed_trades)
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_params = params
|
||||
best_stats = _compute_stats(all_signals, trades)
|
||||
|
||||
if best_params is None:
|
||||
# Every combo scored -inf (too few trades) — still report the
|
||||
# grid's first combination so the fold has *something* to show.
|
||||
best_params = {k: param_grid[k][0] for k in keys}
|
||||
all_signals, trades = _run_combo(candles, precomputed, trade_size, best_params, active_from_index)
|
||||
best_stats = _compute_stats(all_signals, trades)
|
||||
|
||||
return best_params, best_score, best_stats
|
||||
|
||||
|
||||
def _max_drawdown_pct(equity_curve: list[float]) -> float:
|
||||
"""Max peak-to-trough decline of a cumulative-PnL equity curve, as a
|
||||
percentage of the running peak."""
|
||||
if not equity_curve:
|
||||
return 0.0
|
||||
peak = equity_curve[0]
|
||||
max_dd = 0.0
|
||||
for v in equity_curve:
|
||||
peak = max(peak, v)
|
||||
if peak > 0:
|
||||
max_dd = max(max_dd, (peak - v) / peak * 100)
|
||||
return round(max_dd, 2)
|
||||
|
||||
|
||||
async def run_walk_forward(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
exchange: str,
|
||||
timeframe: str = "4h",
|
||||
total_days: int = 1095,
|
||||
train_days: int = 270,
|
||||
test_days: int = 90,
|
||||
trade_size: Decimal = Decimal("10"),
|
||||
param_grid: dict[str, list[float]] | None = None,
|
||||
) -> dict:
|
||||
"""Run a full walk-forward analysis: optimize params per fold on the
|
||||
train window, evaluate out-of-sample on the test window, then stitch
|
||||
all out-of-sample results into one honest performance estimate."""
|
||||
param_grid = param_grid or DEFAULT_PARAM_GRID
|
||||
|
||||
db_symbol = await _fetch_symbol(db, symbol, exchange)
|
||||
if not db_symbol:
|
||||
return {"error": f"Symbol {symbol} not found on {exchange}"}
|
||||
|
||||
folds_spec = generate_folds(total_days, train_days, test_days)
|
||||
if not folds_spec:
|
||||
return {"error": f"total_days ({total_days}) too small for train_days+test_days ({train_days}+{test_days})"}
|
||||
|
||||
fold_results = []
|
||||
stitched_oos_trades: list[dict] = []
|
||||
|
||||
for spec in folds_spec:
|
||||
train_window = await _prepare_window(db, db_symbol.id, timeframe, spec["train_start"], spec["train_end"])
|
||||
if train_window is None:
|
||||
continue
|
||||
train_candles, train_precomputed, train_active_from = train_window
|
||||
|
||||
best_params, _train_score, train_stats = _grid_search(
|
||||
train_candles, train_precomputed, param_grid, trade_size, train_active_from,
|
||||
)
|
||||
|
||||
test_window = await _prepare_window(db, db_symbol.id, timeframe, spec["test_start"], spec["test_end"])
|
||||
if test_window is None:
|
||||
continue
|
||||
test_candles, test_precomputed, test_active_from = test_window
|
||||
|
||||
test_signals, test_trades = _run_combo(test_candles, test_precomputed, trade_size, best_params, test_active_from)
|
||||
test_stats = _compute_stats(test_signals, test_trades)
|
||||
stitched_oos_trades.extend(t for t in test_trades if t.get("status") == "CLOSED")
|
||||
|
||||
fold_results.append({
|
||||
"fold_index": spec["fold_index"],
|
||||
"train_start": spec["train_start"].isoformat(),
|
||||
"train_end": spec["train_end"].isoformat(),
|
||||
"test_start": spec["test_start"].isoformat(),
|
||||
"test_end": spec["test_end"].isoformat(),
|
||||
"best_params": best_params,
|
||||
"in_sample": {
|
||||
"trades": train_stats["trades"]["total"],
|
||||
"win_rate": train_stats["trades"]["win_rate"],
|
||||
"total_pnl": train_stats["trades"]["total_pnl"],
|
||||
"profit_factor": train_stats["trades"]["profit_factor"],
|
||||
},
|
||||
"out_of_sample": {
|
||||
"trades": test_stats["trades"]["total"],
|
||||
"win_rate": test_stats["trades"]["win_rate"],
|
||||
"total_pnl": test_stats["trades"]["total_pnl"],
|
||||
"profit_factor": test_stats["trades"]["profit_factor"],
|
||||
},
|
||||
})
|
||||
|
||||
if not fold_results:
|
||||
return {"error": "No fold had enough candle data to run — try a larger total_days or a lower timeframe."}
|
||||
|
||||
# Stitch OOS trades chronologically — this is the walk-forward's
|
||||
# headline number, the only one that hasn't seen the data it's
|
||||
# evaluated on.
|
||||
stitched_oos_trades.sort(key=lambda t: t["entry_time"])
|
||||
oos_wins = [t for t in stitched_oos_trades if t.get("pnl", 0) > 0]
|
||||
oos_losses = [t for t in stitched_oos_trades if t.get("pnl", 0) <= 0]
|
||||
oos_gross_profit = sum(t.get("pnl", 0) for t in oos_wins)
|
||||
oos_gross_loss = sum(t.get("pnl", 0) for t in oos_losses)
|
||||
oos_total_pnl = sum(t.get("pnl", 0) for t in stitched_oos_trades)
|
||||
oos_win_rate = round(len(oos_wins) / len(stitched_oos_trades) * 100, 1) if stitched_oos_trades else 0.0
|
||||
oos_profit_factor = round(abs(oos_gross_profit / oos_gross_loss), 2) if oos_gross_loss != 0 else None
|
||||
|
||||
equity_curve = [0.0]
|
||||
running = 0.0
|
||||
for t in stitched_oos_trades:
|
||||
running += float(t.get("pnl", 0))
|
||||
equity_curve.append(round(running, 4))
|
||||
|
||||
return {
|
||||
"symbol": symbol,
|
||||
"exchange": exchange,
|
||||
"timeframe": timeframe,
|
||||
"total_days": total_days,
|
||||
"train_days": train_days,
|
||||
"test_days": test_days,
|
||||
"param_grid": param_grid,
|
||||
"folds": fold_results,
|
||||
"out_of_sample_summary": {
|
||||
"trades": len(stitched_oos_trades),
|
||||
"wins": len(oos_wins),
|
||||
"losses": len(oos_losses),
|
||||
"win_rate": oos_win_rate,
|
||||
"total_pnl": round(oos_total_pnl, 2),
|
||||
"profit_factor": oos_profit_factor,
|
||||
"max_drawdown_pct": _max_drawdown_pct(equity_curve),
|
||||
"equity_curve": equity_curve,
|
||||
},
|
||||
}
|
||||
Reference in New Issue
Block a user