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"}
+416
View File
@@ -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,
}
+105 -38
View File
@@ -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(
+303
View File
@@ -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,
},
}