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:
Le
2026-07-04 09:14:32 +07:00
parent 95119b039e
commit 625c2b3773
12 changed files with 1932 additions and 368 deletions
+1 -317
View File
@@ -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
+2
View File
@@ -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)
+160
View File
@@ -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"}