feat: simulate trading fees/slippage in backtest, compute real PnL for real trades

Backtest/walk-forward priced every fill at the exact candle close with zero
cost, making reported win rate/profit factor systematically more optimistic
than live trading. Added configurable taker-fee + slippage simulation
(defaults 0.1%/0.05% per fill) applied to every entry/exit, threaded through
walk-forward's grid search and both API endpoints.

sync_real_trades() hardcoded pnl=0 for every real trade needing it, silently
reporting break-even for real-money trades regardless of actual outcome.
Replaced with FIFO lot matching per (user, symbol, exchange), and fixed
orders.py to persist the exchange's actual average fill price instead of
the (always-None-for-market-orders) requested price, so there's real price
data to match against.

Also verified (and locked in with regression tests) that Divergence/SMC's
pivot-confirmation delay is already causally consistent between live and
backtest — no repaint, no look-ahead leak.

187 backend tests pass (+17).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Le
2026-07-04 14:50:36 +07:00
parent 1bcd15829e
commit 662586c6bc
12 changed files with 733 additions and 90 deletions
+12 -3
View File
@@ -11,7 +11,12 @@ 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.backtest_engine import run_backtest as _run_backtest, MIN_CANDLES
from app.services.backtest_engine import (
run_backtest as _run_backtest,
MIN_CANDLES,
DEFAULT_TAKER_FEE_PCT,
DEFAULT_SLIPPAGE_PCT,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/backtest", tags=["backtest"])
@@ -76,11 +81,13 @@ async def run_backtest(
timeframe: str = Query("30m"),
days: int = Query(7),
trade_size: float = Query(10.0),
fee_pct: float = Query(DEFAULT_TAKER_FEE_PCT, ge=0, description="Round-trip-per-fill taker fee, e.g. 0.001 = 0.1%"),
slippage_pct: float = Query(DEFAULT_SLIPPAGE_PCT, ge=0, description="Adverse slippage per fill, e.g. 0.0005 = 0.05%"),
db: AsyncSession = Depends(get_db),
):
"""Run backtest and return JSON results."""
result = await _run_backtest(
db, symbol, exchange, timeframe, days, Decimal(str(trade_size))
db, symbol, exchange, timeframe, days, Decimal(str(trade_size)), fee_pct, slippage_pct,
)
if "error" in result:
raise HTTPException(status_code=400, detail=result["error"])
@@ -94,7 +101,9 @@ async def run_backtest_post(
timeframe: str = Query("30m"),
days: int = Query(7),
trade_size: float = Query(10.0),
fee_pct: float = Query(DEFAULT_TAKER_FEE_PCT, ge=0),
slippage_pct: float = Query(DEFAULT_SLIPPAGE_PCT, ge=0),
db: AsyncSession = Depends(get_db),
):
"""Alias for GET /backtest/run — supports POST method."""
return await run_backtest(symbol, exchange, timeframe, days, trade_size, db)
return await run_backtest(symbol, exchange, timeframe, days, trade_size, fee_pct, slippage_pct, db)
+6 -2
View File
@@ -96,7 +96,11 @@ async def place_order(
current_user.username, exchange_name, req.symbol, req.side, req.amount,
)
# Persist to real_trades
# Persist to real_trades. For market orders `req.price` is None (no
# limit price was ever set) — the actual execution price only comes
# back from the exchange as `order.average`/`order.price`. Without
# it, this row would have no price at all and PnL could never be
# computed for it later (see sync_real_trades' FIFO PnL matching).
real_trade = RealTrade(
user_id=current_user.id,
exchange=exchange_name,
@@ -104,7 +108,7 @@ async def place_order(
side=req.side,
order_type=req.order_type,
amount=req.amount,
price=req.price,
price=order.average or order.price or req.price,
filled_amount=order.filled,
status=order.status,
order_id=order.order_id,
+4
View File
@@ -21,6 +21,7 @@ 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
from app.services.backtest_engine import DEFAULT_TAKER_FEE_PCT, DEFAULT_SLIPPAGE_PCT
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/walk-forward", tags=["walk_forward"])
@@ -35,6 +36,8 @@ async def run(
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),
fee_pct: float = Query(DEFAULT_TAKER_FEE_PCT, ge=0, description="Round-trip-per-fill taker fee, e.g. 0.001 = 0.1%"),
slippage_pct: float = Query(DEFAULT_SLIPPAGE_PCT, ge=0, description="Adverse slippage per fill, e.g. 0.0005 = 0.05%"),
db: AsyncSession = Depends(get_db),
current_user: UserModel = Depends(get_current_user),
):
@@ -43,6 +46,7 @@ async def run(
db, symbol, exchange, timeframe,
total_days=total_days, train_days=train_days, test_days=test_days,
trade_size=Decimal(str(trade_size)),
fee_pct=fee_pct, slippage_pct=slippage_pct,
)
if "error" in result:
raise HTTPException(status_code=400, detail=result["error"])
+105 -49
View File
@@ -37,6 +37,77 @@ from app.services.signal_scoring import (
# Min candles for warmup: BB(20) + RSI(14) + some room = 30
MIN_CANDLES = 30
# Trading-cost defaults — previously the simulator priced every fill at the
# exact candle close with zero cost, which made every backtest/walk-forward
# report systematically more profitable than live trading could ever be.
# Taker fee: 0.1% per fill (0.2% round-trip) matches typical crypto spot
# taker fees without VIP/token discounts (Binance, MEXC). Slippage: 0.05%
# per fill is a conservative estimate for liquid majors on market orders —
# thinner symbols would see more in reality, but this at least stops
# results from assuming frictionless fills.
DEFAULT_TAKER_FEE_PCT = 0.001
DEFAULT_SLIPPAGE_PCT = 0.0005
def _fill_price(mark_price: float, direction: str, is_entry: bool, slippage_pct: float) -> float:
"""Simulate a realistic market-order fill price — slippage always moves
the price against the trader, never in their favor.
LONG entry / SHORT exit both buy (fill above mark price).
SHORT entry / LONG exit both sell (fill below mark price).
"""
buying = (direction == "LONG") == is_entry
return mark_price * (1 + slippage_pct) if buying else mark_price * (1 - slippage_pct)
def _open_position(direction: str, mark_price: float, timestamp, index: int,
trade_size: Decimal, entry_signal: str,
fee_pct: float, slippage_pct: float) -> dict:
"""Build a new open position dict, applying entry slippage/fees.
`quantity` is sized off the slipped fill price (not the raw mark price)
so that `entry_price * quantity == trade_size`, matching how a real
market order spends a fixed quote-currency amount and receives fewer
units when the fill is worse than the observed price.
"""
entry_price = _fill_price(mark_price, direction, True, slippage_pct)
quantity = float(trade_size) / entry_price
entry_fee = entry_price * quantity * fee_pct
return {
"direction": direction, "entry_price": entry_price,
"entry_time": timestamp, "quantity": quantity,
"entry_signal": entry_signal, "entry_index": index, "status": "OPEN",
"entry_fee": entry_fee,
}
def _close_position(position: dict, mark_price: float, timestamp, exit_reason: str,
fee_pct: float, slippage_pct: float) -> dict:
"""Close an open position in-place, applying exit slippage/fees.
`pnl` is the NET result (gross price movement minus round-trip fees) —
every caller downstream (_compute_stats, walk_forward's fold scoring)
reads `pnl` directly, so this is the only place cost needs to be
subtracted for it to flow through the whole system.
"""
direction = position["direction"]
entry_price = position["entry_price"]
quantity = position["quantity"]
exit_price = _fill_price(mark_price, direction, False, slippage_pct)
exit_fee = exit_price * quantity * fee_pct
if direction == "LONG":
gross_pnl = (exit_price - entry_price) * quantity
else:
gross_pnl = (entry_price - exit_price) * quantity
entry_fee = position.get("entry_fee", 0.0)
fees = entry_fee + exit_fee
position.update({
"exit_price": exit_price, "exit_time": timestamp,
"pnl": gross_pnl - fees, "gross_pnl": gross_pnl, "fees": fees,
"status": "CLOSED", "exit_reason": exit_reason,
})
return position
# Bounded trailing-window sizes used when replaying indicators per candle.
# Every consumer in signal_scoring.py only ever reads the last 1-2 elements
# of these arrays except the BB squeeze check (lookback=10), so windows
@@ -433,12 +504,20 @@ def _simulate_from_scores(
strong_threshold: float = 4.0,
signal_threshold: float = 1.0,
max_hold_candles: int = 48,
fee_pct: float = DEFAULT_TAKER_FEE_PCT,
slippage_pct: float = DEFAULT_SLIPPAGE_PCT,
) -> tuple[list[dict], list[dict]]:
"""Cheap half of simulation: turn a precomputed score series into
signals + trades for one choice of thresholds.
See `_compute_scores_series` for the expensive half — run once,
reused across every threshold combination a grid search tries.
Every fill (entry and exit) goes through `_fill_price`/`_open_position`/
`_close_position` so simulated trades pay the same round-trip fee and
adverse slippage a live market order would, instead of pricing fills at
the exact candle close for free — see DEFAULT_TAKER_FEE_PCT/
DEFAULT_SLIPPAGE_PCT above.
"""
all_signals: list[dict] = []
trades: list[dict] = []
@@ -466,77 +545,43 @@ def _simulate_from_scores(
if signal_type in (STRONG_BUY, BUY):
if current_position and current_position["direction"] == "SHORT":
if signal_type == STRONG_BUY:
entry_price = current_position["entry_price"]
qty = current_position["quantity"]
pnl = (entry_price - latest_close) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
})
_close_position(current_position, latest_close, timestamp, "REVERSAL", fee_pct, slippage_pct)
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",
}
current_position = _open_position(
"LONG", latest_close, timestamp, i, trade_size, signal_type, fee_pct, slippage_pct,
)
elif signal_type in (STRONG_SELL, SELL):
if current_position and current_position["direction"] == "LONG":
if signal_type == STRONG_SELL:
entry_price = current_position["entry_price"]
qty = current_position["quantity"]
pnl = (latest_close - entry_price) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
})
_close_position(current_position, latest_close, timestamp, "REVERSAL", fee_pct, slippage_pct)
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",
}
current_position = _open_position(
"SHORT", latest_close, timestamp, i, trade_size, signal_type, fee_pct, slippage_pct,
)
if current_position and current_position["status"] == "OPEN":
hold = i - current_position["entry_index"]
if hold >= max_hold_candles:
entry_price = current_position["entry_price"]
qty = current_position["quantity"]
if current_position["direction"] == "LONG":
pnl = (latest_close - entry_price) * qty
else:
pnl = (entry_price - latest_close) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "TIME_LIMIT",
})
_close_position(current_position, latest_close, timestamp, "TIME_LIMIT", fee_pct, slippage_pct)
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_price = current_position["entry_price"]
qty = current_position["quantity"]
if current_position["direction"] == "LONG":
pnl = (last_close - entry_price) * qty
else:
pnl = (entry_price - 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",
})
_close_position(
current_position, last_close, candles[-1].timestamp.isoformat(),
"END_OF_DATA", fee_pct, slippage_pct,
)
trades.append(current_position)
return all_signals, trades
@@ -550,6 +595,8 @@ def _simulate_trades(
signal_threshold: float = 1.0,
max_hold_candles: int = 48,
active_from_index: int = MIN_CANDLES,
fee_pct: float = DEFAULT_TAKER_FEE_PCT,
slippage_pct: float = DEFAULT_SLIPPAGE_PCT,
) -> tuple[list[dict], list[dict]]:
"""Convenience wrapper: compute scores then simulate with one set of
thresholds. Callers trying many threshold combinations against the
@@ -558,7 +605,10 @@ def _simulate_trades(
combination instead — see walk_forward.py.
"""
scores_series = _compute_scores_series(candles, precomputed, active_from_index)
return _simulate_from_scores(candles, scores_series, trade_size, strong_threshold, signal_threshold, max_hold_candles)
return _simulate_from_scores(
candles, scores_series, trade_size, strong_threshold, signal_threshold,
max_hold_candles, fee_pct, slippage_pct,
)
def _compute_stats(all_signals: list[dict], trades: list[dict]) -> dict:
@@ -578,6 +628,7 @@ def _compute_stats(all_signals: list[dict], trades: list[dict]) -> dict:
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
total_fees = round(sum(t.get("fees", 0) for t in closed_trades), 2)
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
@@ -595,6 +646,7 @@ def _compute_stats(all_signals: list[dict], trades: list[dict]) -> dict:
"profit_factor": profit_factor,
"avg_win": avg_win,
"avg_loss": avg_loss,
"total_fees": total_fees,
"best_trade": {
"direction": best_trade.get("direction"),
"entry_price": round(best_trade["entry_price"], 4),
@@ -622,6 +674,8 @@ async def run_backtest(
timeframe: str = "30m",
days: int = 7,
trade_size: Decimal = Decimal("10"),
fee_pct: float = DEFAULT_TAKER_FEE_PCT,
slippage_pct: float = DEFAULT_SLIPPAGE_PCT,
) -> dict:
"""Run a single backtest over the last `days` days and return structured results."""
db_symbol = await _fetch_symbol(db, symbol, exchange)
@@ -642,7 +696,9 @@ async def run_backtest(
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)
all_signals, trades = _simulate_trades(
candles, precomputed, trade_size, fee_pct=fee_pct, slippage_pct=slippage_pct,
)
stats = _compute_stats(all_signals, trades)
return {
+117 -28
View File
@@ -8,6 +8,7 @@ from __future__ import annotations
import json
import logging
from collections import defaultdict
from datetime import datetime, timezone, timedelta
from decimal import Decimal
from sqlalchemy import and_, desc, select
@@ -351,55 +352,143 @@ async def execute_signal_trade(
# ═══════════════════════════════════════════════════════════
# Real Trade Sync — closes stale real trades
# Real Trade Sync — closes stale real trades, computes real PnL
# ═══════════════════════════════════════════════════════════
async def sync_real_trades() -> None:
"""Sync real trades: close stale ones, calculate PnL for closed ones.
async def _recompute_realized_pnl(db: AsyncSession) -> int:
"""Recompute realized PnL for every filled real trade via FIFO lot
matching, grouped by (user_id, symbol, exchange).
Called periodically (every 5 min) by the scheduler.
Fixes: real trades were never being closed or having PnL calculated.
`RealTrade` rows represent individual order fills (one buy or one
sell), not paired open/close positions like `HypotheticalTrade` — so a
single row's PnL can't be read off itself. This walks each group's
fills in chronological order, matching each buy against any open SHORT
lots first (then opening a LONG lot with anything left over), and each
sell against open LONG lots first (then opening a SHORT lot) — the
realized PnL from whatever portion closes an existing lot is assigned
to that fill's row. A fill that only opens a new position (nothing to
close yet) has no realized PnL and is left at pnl=None until a later
fill closes it.
Recomputes from the full history every time rather than accumulating
incrementally, which is safe to rerun (idempotent) — the trade volume
this endpoint supports (manual/auto real trading through one connected
exchange per user, not HFT) makes a full recompute cheap enough to run
on every periodic sync.
Rows with no recorded price (e.g. legacy market-order fills persisted
before `orders.py` started storing `order.average`) can't be matched
and are skipped — they neither consume nor produce a lot, so a FIFO
walk that includes them would silently misattribute PnL to whichever
fill happens to match next.
"""
result = await db.execute(
select(RealTrade)
.where(RealTrade.filled_amount > 0, RealTrade.price.isnot(None))
.order_by(RealTrade.user_id, RealTrade.symbol, RealTrade.exchange, RealTrade.created_at.asc())
)
trades = result.scalars().all()
groups: dict[tuple, list[RealTrade]] = defaultdict(list)
for t in trades:
groups[(t.user_id, t.symbol, t.exchange)].append(t)
updated = 0
for _key, group in groups.items():
long_lots: list[list] = [] # each lot: [remaining_qty: Decimal, entry_price: Decimal]
short_lots: list[list] = []
for t in group:
side = (t.side or "").lower()
qty = Decimal(t.filled_amount)
price = Decimal(t.price)
if side not in ("buy", "sell") or qty <= 0 or price <= 0:
continue
opposing = short_lots if side == "buy" else long_lots
same_side_lots = long_lots if side == "buy" else short_lots
remaining = qty
realized = Decimal("0")
matched_cost = Decimal("0")
while remaining > 0 and opposing:
lot = opposing[0]
matched = min(remaining, lot[0])
if side == "buy":
realized += (lot[1] - price) * matched # closing a short: profit if price fell
else:
realized += (price - lot[1]) * matched # closing a long: profit if price rose
matched_cost += lot[1] * matched
lot[0] -= matched
remaining -= matched
if lot[0] <= 0:
opposing.pop(0)
if remaining > 0:
same_side_lots.append([remaining, price])
if matched_cost > 0:
pnl_percent = (realized / matched_cost) * Decimal("100")
if t.pnl != realized or t.pnl_percent != pnl_percent:
t.pnl = realized
t.pnl_percent = pnl_percent
updated += 1
elif t.pnl is not None:
# Pure opening fill — nothing existed yet to realize against.
t.pnl = None
t.pnl_percent = None
updated += 1
return updated
async def sync_real_trades() -> None:
"""Sync real trades: close stale open orders, then (re)compute realized
PnL for every filled trade via FIFO matching (see
`_recompute_realized_pnl`).
Called periodically (every 5 min) by the scheduler. Previously this
force-set pnl=0 for every stale-closed AND every already-closed-but-
unpriced trade regardless of whether it actually filled — silently
reporting "break-even" for real-money trades that may have made or
lost real money. Now pnl is only ever 0/None for orders that genuinely
never filled; any order with a nonzero fill gets a real FIFO-matched
PnL once a later trade closes its position.
"""
async with async_session_factory() as db:
# 1. Fetch open real trades
# 1. Force-close real orders stuck "open" too long (exchange likely
# never filled them). Only the STATUS changes here for anything
# with a partial fill — its PnL (if any) comes from the FIFO
# recompute below once/if a later trade closes it out.
result = await db.execute(
select(RealTrade).where(RealTrade.status == "open")
)
open_trades = result.scalars().all()
if not open_trades:
return
now = datetime.now(timezone.utc)
stale_closed = 0
for trade in open_trades:
hold_duration = now - trade.created_at
if hold_duration > timedelta(hours=24):
trade.status = "closed"
trade.closed_at = now
trade.pnl = Decimal("0")
trade.pnl_percent = Decimal("0")
if not trade.filled_amount or trade.filled_amount <= 0:
# Never filled at all — there is genuinely no PnL.
trade.pnl = Decimal("0")
trade.pnl_percent = Decimal("0")
stale_closed += 1
logger.info(
"🔒 Real trade #%d CLOSED (time limit 24h): %s %s",
trade.id, trade.side, trade.symbol,
)
# 2. Calculate PnL for closed trades missing it
result2 = await db.execute(
select(RealTrade).where(
and_(RealTrade.status.in_(["closed", "filled"]),
RealTrade.pnl.is_(None))
)
)
closed_no_pnl = result2.scalars().all()
for trade in closed_no_pnl:
trade.pnl = Decimal("0")
trade.pnl_percent = Decimal("0")
# 2. Recompute realized PnL for every filled trade (not just ones
# missing it — a new closing fill can retroactively realize PnL on
# an earlier opening fill that already had pnl=None).
updated = await _recompute_realized_pnl(db)
await db.commit()
if open_trades or closed_no_pnl:
if stale_closed or updated:
logger.info(
"Real trade sync: %d open checked, %d closed PnL fixed",
len(open_trades), len(closed_no_pnl),
"Real trade sync: %d stale trades closed, %d PnL rows recomputed",
stale_closed, updated,
)
+25 -6
View File
@@ -26,6 +26,8 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.services.backtest_engine import (
MIN_CANDLES,
DEFAULT_TAKER_FEE_PCT,
DEFAULT_SLIPPAGE_PCT,
_fetch_symbol,
_fetch_candles,
_precompute_indicators,
@@ -138,7 +140,7 @@ async def _prepare_window(
return candles, precomputed, active_from_index
def _run_combo(candles, scores_series, trade_size, params):
def _run_combo(candles, scores_series, trade_size, params, fee_pct, slippage_pct):
"""Apply one parameter combination to an already-computed score series
(see `_compute_scores_series`) — cheap, no indicator recomputation."""
return _simulate_from_scores(
@@ -146,6 +148,8 @@ def _run_combo(candles, scores_series, trade_size, params):
strong_threshold=params["strong_threshold"],
signal_threshold=params["signal_threshold"],
max_hold_candles=int(params["max_hold_candles"]),
fee_pct=fee_pct,
slippage_pct=slippage_pct,
)
@@ -154,12 +158,20 @@ def _grid_search(
scores_series: list[dict],
param_grid: dict[str, list[float]],
trade_size: Decimal,
fee_pct: float = DEFAULT_TAKER_FEE_PCT,
slippage_pct: float = DEFAULT_SLIPPAGE_PCT,
) -> tuple[dict[str, float], float, dict]:
"""Try every combination in param_grid against one precomputed score
series, return the one that scores best on the train window by
`_fold_score`. The 13-algorithm scoring itself already happened once
to produce `scores_series` — this loop only replays cheap threshold
comparisons and trade bookkeeping per combination."""
comparisons and trade bookkeeping per combination.
Fees/slippage are applied here too (not just the final OOS report) so
the grid search picks parameters that hold up after trading costs,
instead of favoring high-frequency combos that only look good when
fills are free.
"""
keys = list(param_grid.keys())
best_params: dict[str, float] | None = None
best_score = float("-inf")
@@ -167,7 +179,7 @@ def _grid_search(
for combo in product(*(param_grid[k] for k in keys)):
params = dict(zip(keys, combo))
all_signals, trades = _run_combo(candles, scores_series, trade_size, params)
all_signals, trades = _run_combo(candles, scores_series, trade_size, params, fee_pct, slippage_pct)
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
score = _fold_score(closed_trades)
if score > best_score:
@@ -179,7 +191,7 @@ def _grid_search(
# 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, scores_series, trade_size, best_params)
all_signals, trades = _run_combo(candles, scores_series, trade_size, best_params, fee_pct, slippage_pct)
best_stats = _compute_stats(all_signals, trades)
return best_params, best_score, best_stats
@@ -209,6 +221,8 @@ async def run_walk_forward(
test_days: int = 90,
trade_size: Decimal = Decimal("10"),
param_grid: dict[str, list[float]] | None = None,
fee_pct: float = DEFAULT_TAKER_FEE_PCT,
slippage_pct: float = DEFAULT_SLIPPAGE_PCT,
) -> 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
@@ -234,7 +248,7 @@ async def run_walk_forward(
train_scores = _compute_scores_series(train_candles, train_precomputed, train_active_from)
best_params, _train_score, train_stats = _grid_search(
train_candles, train_scores, param_grid, trade_size,
train_candles, train_scores, param_grid, trade_size, fee_pct, slippage_pct,
)
test_window = await _prepare_window(db, db_symbol.id, timeframe, spec["test_start"], spec["test_end"])
@@ -243,7 +257,7 @@ async def run_walk_forward(
test_candles, test_precomputed, test_active_from = test_window
test_scores = _compute_scores_series(test_candles, test_precomputed, test_active_from)
test_signals, test_trades = _run_combo(test_candles, test_scores, trade_size, best_params)
test_signals, test_trades = _run_combo(test_candles, test_scores, trade_size, best_params, fee_pct, slippage_pct)
test_stats = _compute_stats(test_signals, test_trades)
stitched_oos_trades.extend(t for t in test_trades if t.get("status") == "CLOSED")
@@ -265,6 +279,7 @@ async def run_walk_forward(
"win_rate": test_stats["trades"]["win_rate"],
"total_pnl": test_stats["trades"]["total_pnl"],
"profit_factor": test_stats["trades"]["profit_factor"],
"total_fees": test_stats["trades"]["total_fees"],
},
})
@@ -282,6 +297,7 @@ async def run_walk_forward(
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
oos_total_fees = round(sum(t.get("fees", 0) for t in stitched_oos_trades), 2)
equity_curve = [0.0]
running = 0.0
@@ -297,6 +313,8 @@ async def run_walk_forward(
"train_days": train_days,
"test_days": test_days,
"param_grid": param_grid,
"fee_pct": fee_pct,
"slippage_pct": slippage_pct,
"folds": fold_results,
"out_of_sample_summary": {
"trades": len(stitched_oos_trades),
@@ -305,6 +323,7 @@ async def run_walk_forward(
"win_rate": oos_win_rate,
"total_pnl": round(oos_total_pnl, 2),
"profit_factor": oos_profit_factor,
"total_fees": oos_total_fees,
"max_drawdown_pct": _max_drawdown_pct(equity_curve),
"equity_curve": equity_curve,
},