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
+192
View File
@@ -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 == []
+193
View File
@@ -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