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:
@@ -0,0 +1,192 @@
|
||||
"""Tests for the backtest engine refactor (fetch/precompute/simulate split).
|
||||
|
||||
Focused on: the new fetch helpers behave correctly against the DB, and
|
||||
`_simulate_trades` correctly threads its threshold/max-hold parameters
|
||||
through to classification and exit logic — the whole reason it was split
|
||||
out of `_run_backtest` was so walk_forward.py could vary these cheaply.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from app.models.candle import Candle
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
from app.services import backtest_engine
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
async def _seed_symbol(db_session, name="BTC/USDT", exchange_name="mexc"):
|
||||
exchange = Exchange(name=exchange_name, display_name=exchange_name.upper())
|
||||
db_session.add(exchange)
|
||||
await db_session.flush()
|
||||
symbol = Symbol(symbol=name, base=name.split("/")[0], quote=name.split("/")[1], exchange_id=exchange.id)
|
||||
db_session.add(symbol)
|
||||
await db_session.flush()
|
||||
return exchange, symbol
|
||||
|
||||
|
||||
async def _seed_candles(db_session, symbol_id, timeframe, start, count, step, price_fn):
|
||||
for i in range(count):
|
||||
ts = start + step * i
|
||||
price = price_fn(i)
|
||||
db_session.add(Candle(
|
||||
symbol_id=symbol_id, timeframe=timeframe, timestamp=ts,
|
||||
open=Decimal(str(price)), high=Decimal(str(price * 1.01)),
|
||||
low=Decimal(str(price * 0.99)), close=Decimal(str(price)),
|
||||
volume=Decimal("1000"),
|
||||
))
|
||||
await db_session.flush()
|
||||
|
||||
|
||||
async def test_fetch_symbol_found_and_not_found(db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
found = await backtest_engine._fetch_symbol(db_session, "BTC/USDT", "mexc")
|
||||
assert found is not None
|
||||
assert found.id == symbol.id
|
||||
|
||||
missing = await backtest_engine._fetch_symbol(db_session, "ETH/USDT", "mexc")
|
||||
assert missing is None
|
||||
|
||||
|
||||
async def test_fetch_candles_respects_since_and_until(db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
base = datetime(2026, 1, 1, tzinfo=timezone.utc)
|
||||
await _seed_candles(db_session, symbol.id, "1h", base, 10, timedelta(hours=1), lambda i: 100 + i)
|
||||
|
||||
# Full range
|
||||
all_candles = await backtest_engine._fetch_candles(db_session, symbol.id, "1h", since=base)
|
||||
assert len(all_candles) == 10
|
||||
|
||||
# since excludes earlier candles
|
||||
later = await backtest_engine._fetch_candles(db_session, symbol.id, "1h", since=base + timedelta(hours=5))
|
||||
assert len(later) == 5
|
||||
|
||||
# until excludes candles at/after the boundary
|
||||
earlier = await backtest_engine._fetch_candles(db_session, symbol.id, "1h", since=base, until=base + timedelta(hours=5))
|
||||
assert len(earlier) == 5
|
||||
assert all(c.timestamp < base + timedelta(hours=5) for c in earlier)
|
||||
|
||||
|
||||
async def test_run_backtest_symbol_not_found(db_session):
|
||||
result = await backtest_engine.run_backtest(db_session, "DOES/NOTEXIST", "mexc")
|
||||
assert "error" in result
|
||||
assert "not found" in result["error"]
|
||||
|
||||
|
||||
async def test_run_backtest_insufficient_candles(db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
base = datetime.now(timezone.utc) - timedelta(hours=5)
|
||||
await _seed_candles(db_session, symbol.id, "1h", base, 5, timedelta(hours=1), lambda i: 100)
|
||||
|
||||
result = await backtest_engine.run_backtest(db_session, "BTC/USDT", "mexc", timeframe="1h", days=1)
|
||||
assert "error" in result
|
||||
assert "at least" in result["error"]
|
||||
|
||||
|
||||
async def test_run_backtest_returns_expected_shape(db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
base = datetime.now(timezone.utc) - timedelta(hours=60)
|
||||
# Gentle random-ish walk — not asserting on trade content, just structure.
|
||||
await _seed_candles(db_session, symbol.id, "1h", base, 60, timedelta(hours=1),
|
||||
lambda i: 100 + 5 * math.sin(i / 3))
|
||||
|
||||
result = await backtest_engine.run_backtest(db_session, "BTC/USDT", "mexc", timeframe="1h", days=3)
|
||||
assert "error" not in result
|
||||
assert result["symbol"] == "BTC/USDT"
|
||||
assert result["candles_count"] == 60
|
||||
assert "trades" in result
|
||||
assert set(result["trades"].keys()) >= {"total", "wins", "losses", "win_rate", "total_pnl", "profit_factor"}
|
||||
|
||||
|
||||
def _make_fake_classifier(buy_at: set[int], sell_at: set[int]):
|
||||
"""Build a fake `_classify_signal_combined` that ignores indicator data
|
||||
and instead signals BUY/SELL purely from `close_price`, which the real
|
||||
test encodes as the candle index (so we can trigger deterministically).
|
||||
Also records the strong_threshold/signal_threshold it was called with.
|
||||
"""
|
||||
calls = []
|
||||
|
||||
def fake(close_price, *args, **kwargs):
|
||||
calls.append({
|
||||
"strong_threshold": kwargs.get("strong_threshold"),
|
||||
"signal_threshold": kwargs.get("signal_threshold"),
|
||||
})
|
||||
idx = int(round(close_price))
|
||||
if idx in buy_at:
|
||||
return backtest_engine.STRONG_BUY, "STRONG", 0.9, {}
|
||||
if idx in sell_at:
|
||||
return backtest_engine.STRONG_SELL, "STRONG", 0.9, {}
|
||||
return None, None, 0.5, {}
|
||||
|
||||
return fake, calls
|
||||
|
||||
|
||||
async def test_simulate_trades_passes_thresholds_to_classifier(monkeypatch, db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
base = datetime.now(timezone.utc) - timedelta(hours=40)
|
||||
# close price == candle index, so the fake classifier can key off it
|
||||
await _seed_candles(db_session, symbol.id, "1h", base, 40, timedelta(hours=1), lambda i: i)
|
||||
candles = await backtest_engine._fetch_candles(db_session, symbol.id, "1h", since=base)
|
||||
precomputed = backtest_engine._precompute_indicators(candles, "1h")
|
||||
|
||||
fake, calls = _make_fake_classifier(buy_at={35}, sell_at=set())
|
||||
monkeypatch.setattr(backtest_engine, "_classify_signal_combined", fake)
|
||||
|
||||
backtest_engine._simulate_trades(
|
||||
candles, precomputed, Decimal("10"),
|
||||
strong_threshold=2.5, signal_threshold=0.5,
|
||||
)
|
||||
|
||||
assert len(calls) > 0
|
||||
assert all(c["strong_threshold"] == 2.5 for c in calls)
|
||||
assert all(c["signal_threshold"] == 0.5 for c in calls)
|
||||
|
||||
|
||||
async def test_simulate_trades_exits_on_max_hold_candles(monkeypatch, db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
base = datetime.now(timezone.utc) - timedelta(hours=40)
|
||||
await _seed_candles(db_session, symbol.id, "1h", base, 40, timedelta(hours=1), lambda i: i)
|
||||
candles = await backtest_engine._fetch_candles(db_session, symbol.id, "1h", since=base)
|
||||
precomputed = backtest_engine._precompute_indicators(candles, "1h")
|
||||
|
||||
# Open a LONG at index 30 (candle price 30) and never signal again —
|
||||
# it must be force-closed exactly `max_hold_candles` candles later.
|
||||
fake, _ = _make_fake_classifier(buy_at={30}, sell_at=set())
|
||||
monkeypatch.setattr(backtest_engine, "_classify_signal_combined", fake)
|
||||
|
||||
_, trades = backtest_engine._simulate_trades(
|
||||
candles, precomputed, Decimal("10"), max_hold_candles=5,
|
||||
)
|
||||
|
||||
assert len(trades) == 1
|
||||
trade = trades[0]
|
||||
assert trade["status"] == "CLOSED"
|
||||
assert trade["exit_reason"] == "TIME_LIMIT"
|
||||
expected_exit_index = trade["entry_index"] + 5
|
||||
assert trade["exit_time"] == candles[expected_exit_index].timestamp.isoformat()
|
||||
|
||||
|
||||
async def test_simulate_trades_active_from_index_skips_warmup_region(monkeypatch, db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
base = datetime.now(timezone.utc) - timedelta(hours=40)
|
||||
await _seed_candles(db_session, symbol.id, "1h", base, 40, timedelta(hours=1), lambda i: i)
|
||||
candles = await backtest_engine._fetch_candles(db_session, symbol.id, "1h", since=base)
|
||||
precomputed = backtest_engine._precompute_indicators(candles, "1h")
|
||||
|
||||
# A BUY signal planted inside the warmup region (index 32, but
|
||||
# active_from_index=35) must never open a trade.
|
||||
fake, _ = _make_fake_classifier(buy_at={32}, sell_at=set())
|
||||
monkeypatch.setattr(backtest_engine, "_classify_signal_combined", fake)
|
||||
|
||||
all_signals, trades = backtest_engine._simulate_trades(
|
||||
candles, precomputed, Decimal("10"), active_from_index=35,
|
||||
)
|
||||
|
||||
assert trades == []
|
||||
assert all_signals == []
|
||||
@@ -0,0 +1,193 @@
|
||||
"""Tests for the walk-forward optimizer (item m — overfitting mitigation).
|
||||
|
||||
Covers the pure planning/scoring functions directly, and runs one small
|
||||
end-to-end pass against seeded synthetic candle data to prove the fold
|
||||
loop, grid search, and out-of-sample stitching all wire together
|
||||
correctly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from app.models.candle import Candle
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
from app.services import walk_forward
|
||||
|
||||
|
||||
# ── generate_folds ──────────────────────────────────────────────────────
|
||||
|
||||
def test_generate_folds_basic_counts_and_boundaries():
|
||||
anchor = datetime(2026, 1, 1, tzinfo=timezone.utc)
|
||||
folds = walk_forward.generate_folds(total_days=365, train_days=180, test_days=60, anchor=anchor)
|
||||
|
||||
# (365 - 180) / 60 = 3.08 -> folds while test_end <= anchor
|
||||
assert len(folds) >= 1
|
||||
for f in folds:
|
||||
assert f["train_end"] == f["test_start"]
|
||||
assert (f["train_end"] - f["train_start"]).days == 180
|
||||
assert (f["test_end"] - f["test_start"]).days == 60
|
||||
assert f["test_end"] <= anchor
|
||||
|
||||
# Folds walk forward chronologically, each starting test_days later
|
||||
for a, b in zip(folds, folds[1:]):
|
||||
assert b["train_start"] - a["train_start"] == timedelta(days=60)
|
||||
|
||||
|
||||
def test_generate_folds_too_small_returns_empty():
|
||||
folds = walk_forward.generate_folds(total_days=100, train_days=180, test_days=60)
|
||||
assert folds == []
|
||||
|
||||
|
||||
# ── _fold_score ──────────────────────────────────────────────────────────
|
||||
|
||||
def test_fold_score_rejects_too_few_trades():
|
||||
trades = [{"pnl": 5.0}, {"pnl": 3.0}] # below MIN_TRADES_PER_FOLD
|
||||
assert walk_forward._fold_score(trades) == float("-inf")
|
||||
|
||||
|
||||
def test_fold_score_prefers_consistent_edge_over_lucky_streak():
|
||||
# Same total PnL (50), but one is steady small wins, the other is one
|
||||
# huge win plus several losses — the steadier one should score higher.
|
||||
consistent = [{"pnl": 10.0} for _ in range(5)]
|
||||
lucky = [{"pnl": 50.0}, {"pnl": -10.0}, {"pnl": -10.0}, {"pnl": -10.0}, {"pnl": -10.0}]
|
||||
|
||||
score_consistent = walk_forward._fold_score(consistent)
|
||||
score_lucky = walk_forward._fold_score(lucky)
|
||||
|
||||
assert score_consistent > score_lucky
|
||||
|
||||
|
||||
def test_fold_score_zero_variance_all_same_sign():
|
||||
trades = [{"pnl": 10.0} for _ in range(6)]
|
||||
score = walk_forward._fold_score(trades)
|
||||
assert score == pytest.approx(10.0 * math.sqrt(6))
|
||||
|
||||
|
||||
# ── _max_drawdown_pct ──────────────────────────────────────────────────
|
||||
|
||||
def test_max_drawdown_pct_known_curve():
|
||||
# Peak at 100, trough at 80 -> 20% drawdown
|
||||
curve = [0, 50, 100, 80, 90, 120]
|
||||
assert walk_forward._max_drawdown_pct(curve) == pytest.approx(20.0)
|
||||
|
||||
|
||||
def test_max_drawdown_pct_empty_curve():
|
||||
assert walk_forward._max_drawdown_pct([]) == 0.0
|
||||
|
||||
|
||||
# ── _grid_search ─────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_grid_search_picks_the_best_scoring_combo(monkeypatch):
|
||||
"""Fake `_run_combo` so each parameter combo deterministically returns
|
||||
a trade set with a known score, then assert grid search picks the
|
||||
best one and reports its stats."""
|
||||
good_params = {"strong_threshold": 4.5, "signal_threshold": 1.5, "max_hold_candles": 96}
|
||||
|
||||
def fake_run_combo(candles, precomputed, trade_size, params, active_from_index):
|
||||
if params == good_params:
|
||||
trades = [{"pnl": 10.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 110.0} for _ in range(10)]
|
||||
else:
|
||||
trades = [
|
||||
{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0},
|
||||
{"pnl": -1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 99.0},
|
||||
]
|
||||
return [], trades
|
||||
|
||||
monkeypatch.setattr(walk_forward, "_run_combo", fake_run_combo)
|
||||
|
||||
grid = {"strong_threshold": [3.5, 4.5], "signal_threshold": [1.0, 1.5], "max_hold_candles": [48, 96]}
|
||||
best_params, best_score, best_stats = walk_forward._grid_search(
|
||||
candles=[], precomputed={}, param_grid=grid, trade_size=Decimal("10"), active_from_index=0,
|
||||
)
|
||||
|
||||
assert best_params == good_params
|
||||
assert best_stats["trades"]["total"] == 10
|
||||
assert best_score > float("-inf")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_grid_search_falls_back_when_every_combo_too_sparse(monkeypatch):
|
||||
def fake_run_combo(candles, precomputed, trade_size, params, active_from_index):
|
||||
# 1 trade, below MIN_TRADES_PER_FOLD
|
||||
return [], [{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0}]
|
||||
|
||||
monkeypatch.setattr(walk_forward, "_run_combo", fake_run_combo)
|
||||
|
||||
grid = {"strong_threshold": [3.5, 4.5], "signal_threshold": [1.0], "max_hold_candles": [48]}
|
||||
best_params, best_score, best_stats = walk_forward._grid_search(
|
||||
candles=[], precomputed={}, param_grid=grid, trade_size=Decimal("10"), active_from_index=0,
|
||||
)
|
||||
|
||||
# Falls back to the first grid combination rather than raising
|
||||
assert best_params == {"strong_threshold": 3.5, "signal_threshold": 1.0, "max_hold_candles": 48}
|
||||
assert best_score == float("-inf")
|
||||
assert best_stats["trades"]["total"] == 1
|
||||
|
||||
|
||||
# ── end-to-end (small synthetic dataset) ─────────────────────────────────
|
||||
|
||||
async def _seed_symbol(db_session, name="BTC/USDT", exchange_name="mexc"):
|
||||
exchange = Exchange(name=exchange_name, display_name=exchange_name.upper())
|
||||
db_session.add(exchange)
|
||||
await db_session.flush()
|
||||
symbol = Symbol(symbol=name, base=name.split("/")[0], quote=name.split("/")[1], exchange_id=exchange.id)
|
||||
db_session.add(symbol)
|
||||
await db_session.flush()
|
||||
return exchange, symbol
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_walk_forward_end_to_end_on_synthetic_data(db_session):
|
||||
_, symbol = await _seed_symbol(db_session)
|
||||
|
||||
anchor = datetime.now(timezone.utc)
|
||||
total_days = 40
|
||||
train_days = 20
|
||||
test_days = 10
|
||||
# 4h candles over 40 days = 240 candles — small enough to run fast,
|
||||
# oscillating so the 13-algorithm system has *something* to react to.
|
||||
start = anchor - timedelta(days=total_days + 5) # + warmup headroom
|
||||
count = int((total_days + 5) * 24 / 4)
|
||||
for i in range(count):
|
||||
price = 100 + 10 * math.sin(i / 5) + (i % 7)
|
||||
ts = start + timedelta(hours=4 * i)
|
||||
db_session.add(Candle(
|
||||
symbol_id=symbol.id, timeframe="4h", timestamp=ts,
|
||||
open=Decimal(str(price)), high=Decimal(str(price * 1.02)),
|
||||
low=Decimal(str(price * 0.98)), close=Decimal(str(price)),
|
||||
volume=Decimal("1000"),
|
||||
))
|
||||
await db_session.flush()
|
||||
|
||||
small_grid = {"strong_threshold": [4.0], "signal_threshold": [1.0], "max_hold_candles": [48]}
|
||||
result = await walk_forward.run_walk_forward(
|
||||
db_session, "BTC/USDT", "mexc", timeframe="4h",
|
||||
total_days=total_days, train_days=train_days, test_days=test_days,
|
||||
param_grid=small_grid,
|
||||
)
|
||||
|
||||
assert "error" not in result
|
||||
expected_fold_count = len(walk_forward.generate_folds(total_days, train_days, test_days, anchor=anchor))
|
||||
assert len(result["folds"]) <= expected_fold_count
|
||||
assert len(result["folds"]) >= 1
|
||||
|
||||
summary = result["out_of_sample_summary"]
|
||||
assert set(summary.keys()) >= {"trades", "win_rate", "total_pnl", "profit_factor", "max_drawdown_pct", "equity_curve"}
|
||||
assert summary["equity_curve"][0] == 0.0
|
||||
assert len(summary["equity_curve"]) == summary["trades"] + 1
|
||||
|
||||
for fold in result["folds"]:
|
||||
assert set(fold["best_params"].keys()) == {"strong_threshold", "signal_threshold", "max_hold_candles"}
|
||||
assert "in_sample" in fold and "out_of_sample" in fold
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_walk_forward_symbol_not_found(db_session):
|
||||
result = await walk_forward.run_walk_forward(db_session, "NOPE/USDT", "mexc")
|
||||
assert "error" in result
|
||||
Reference in New Issue
Block a user