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
@@ -0,0 +1,52 @@
"""Add walk_forward_results table
Stores walk-forward backtest runs: rolling train/test folds with
per-fold optimized parameters and in-sample vs out-of-sample metrics,
plus the aggregated out-of-sample summary. `result_json` holds the
full fold-by-fold detail; the flat columns are for fast history listing.
Revision ID: add_walk_forward_results
Revises: merge_heads_1
Create Date: 2026-07-04
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = "add_walk_forward_results"
down_revision: Union[str, None] = "merge_heads_1"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
op.create_table(
"walk_forward_results",
sa.Column("id", sa.UUID(), nullable=False),
sa.Column("user_id", sa.UUID(), nullable=False),
sa.Column("symbol", sa.String(length=50), nullable=False),
sa.Column("exchange", sa.String(length=20), nullable=False),
sa.Column("timeframe", sa.String(length=10), nullable=False),
sa.Column("total_days", sa.Integer(), nullable=False),
sa.Column("train_days", sa.Integer(), nullable=False),
sa.Column("test_days", sa.Integer(), nullable=False),
sa.Column("folds_count", sa.Integer(), nullable=False),
sa.Column("oos_trades", sa.Integer(), nullable=False),
sa.Column("oos_win_rate", sa.Numeric(precision=6, scale=2), nullable=True),
sa.Column("oos_total_pnl", sa.Numeric(precision=20, scale=8), nullable=True),
sa.Column("oos_profit_factor", sa.Numeric(precision=10, scale=4), nullable=True),
sa.Column("oos_max_drawdown_pct", sa.Numeric(precision=6, scale=2), nullable=True),
sa.Column("result_json", sa.Text(), nullable=False, comment="Full fold-by-fold detail + stitched OOS equity curve"),
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
sa.ForeignKeyConstraint(["user_id"], ["users.id"]),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_walk_forward_results_user_id", "walk_forward_results", ["user_id"], unique=False)
op.create_index("ix_walk_forward_results_created_at", "walk_forward_results", ["created_at"], unique=False)
def downgrade() -> None:
op.drop_index("ix_walk_forward_results_created_at", table_name="walk_forward_results")
op.drop_index("ix_walk_forward_results_user_id", table_name="walk_forward_results")
op.drop_table("walk_forward_results")
+24
View File
@@ -0,0 +1,24 @@
"""merge divergent heads (1b3f1630986f, 4_add_sl_tp_columns)
Both branched off add_candle_partitions independently, leaving two
unmerged heads. This is a no-op merge so `alembic upgrade head` has a
single target again.
Revision ID: merge_heads_1
Revises: 1b3f1630986f, 4_add_sl_tp_columns
Create Date: 2026-07-04
"""
from typing import Sequence, Union
revision: str = "merge_heads_1"
down_revision: Union[str, Sequence[str], None] = ("1b3f1630986f", "4_add_sl_tp_columns")
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
pass
def downgrade() -> None:
pass
+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,
},
}
+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
+364 -13
View File
@@ -1,6 +1,6 @@
import { useState, useCallback, useEffect } from 'react';
import { useT } from '../../translations';
import { apiFetch } from '../api/apiService';
import { apiFetch, ApiServiceError } from '../api/apiService';
interface SignalCounts {
[key: string]: number;
@@ -50,6 +50,65 @@ interface BacktestResult {
};
}
// ── Walk-forward types ──
interface WfBestParams {
strong_threshold: number;
signal_threshold: number;
max_hold_candles: number;
}
interface WfFoldMetrics {
trades: number;
win_rate: number;
total_pnl: number;
profit_factor: number | null;
}
interface WfFoldResult {
fold_index: number;
train_start: string;
train_end: string;
test_start: string;
test_end: string;
best_params: WfBestParams;
in_sample: WfFoldMetrics;
out_of_sample: WfFoldMetrics;
}
interface WalkForwardResult {
id?: string;
symbol: string;
exchange: string;
timeframe: string;
total_days: number;
train_days: number;
test_days: number;
folds: WfFoldResult[];
out_of_sample_summary: {
trades: number;
wins: number;
losses: number;
win_rate: number;
total_pnl: number;
profit_factor: number | null;
max_drawdown_pct: number;
equity_curve: number[];
};
}
interface WfHistoryItem {
id: string;
symbol: string;
exchange: string;
timeframe: string;
total_days: number;
train_days: number;
test_days: number;
folds_count: number;
oos_trades: number;
oos_win_rate: number | null;
oos_total_pnl: number | null;
oos_profit_factor: number | null;
oos_max_drawdown_pct: number | null;
created_at: string;
}
const SIGNAL_ICONS: Record<string, string> = {
'STRONG_BUY': '🚀', 'BUY': '📈',
'STRONG_SELL': '🔻', 'SELL': '📉',
@@ -59,25 +118,23 @@ const SIGNAL_ICONS: Record<string, string> = {
const EXCHANGES = ['binance', 'bybit', 'mexc', 'gate', 'bingx'];
const TIMEFRAMES = ['15m', '30m', '1h', '4h'];
const WF_TIMEFRAMES = ['1h', '4h'];
const selectClass = 'rounded-md border border-border-default bg-bg-surface px-2.5 py-1.5 text-sm text-text-primary';
const thClass = 'px-3 py-2 text-left font-medium text-text-secondary';
const tdClass = 'px-3 py-1.5';
const cardClass = 'rounded-lg bg-bg-surface p-4';
const modeBtnClass = 'min-h-9 touch-manipulation rounded-md border px-4 py-1.5 text-sm';
export default function BacktestPage() {
const { t } = useT();
const [mode, setMode] = useState<'single' | 'walk-forward'>('single');
const [exchange, setExchange] = useState('binance');
const [symbol, setSymbol] = useState('BTC/USDT');
const [symbols, setSymbols] = useState<string[]>([]);
const [timeframe, setTimeframe] = useState('30m');
const [days, setDays] = useState(7);
const [tradeSize, setTradeSize] = useState(10);
const [result, setResult] = useState<BacktestResult | null>(null);
const [loading, setLoading] = useState(false);
const [error, setError] = useState('');
// Load symbols when exchange changes
// Load symbols when exchange changes — shared between both modes
useEffect(() => {
let cancelled = false;
async function loadSymbols() {
@@ -86,7 +143,6 @@ export default function BacktestPage() {
if (cancelled) return;
const names = (data.symbols || []).map((s: any) => s.symbol);
setSymbols(names);
// Keep current symbol if in list, else pick first
if (names.length > 0 && !names.includes(symbol)) {
setSymbol(names[0]);
}
@@ -96,6 +152,43 @@ export default function BacktestPage() {
return () => { cancelled = true; };
}, [exchange]);
return (
<div className="mx-auto max-w-[1000px] bg-bg-primary p-5 text-text-primary">
<h1 className="text-2xl font-bold text-text-heading">📊 {t('Backtest')}</h1>
<div className="mb-5 mt-3 flex gap-1">
<button onClick={() => setMode('single')} className={`${modeBtnClass} ${mode === 'single' ? 'border-accent-blue bg-accent-blue font-bold text-white' : 'border-border-default bg-bg-hover text-text-secondary'}`}>
▶ Single Run
</button>
<button onClick={() => setMode('walk-forward')} className={`${modeBtnClass} ${mode === 'walk-forward' ? 'border-accent-blue bg-accent-blue font-bold text-white' : 'border-border-default bg-bg-hover text-text-secondary'}`}>
🧪 Walk-Forward
</button>
</div>
{mode === 'single' ? (
<SingleRunView exchange={exchange} setExchange={setExchange} symbol={symbol} setSymbol={setSymbol}
symbols={symbols} tradeSize={tradeSize} setTradeSize={setTradeSize} />
) : (
<WalkForwardView exchange={exchange} setExchange={setExchange} symbol={symbol} setSymbol={setSymbol}
symbols={symbols} tradeSize={tradeSize} setTradeSize={setTradeSize} />
)}
</div>
);
}
// ═══════════════ SINGLE RUN ═══════════════
function SingleRunView({ exchange, setExchange, symbol, setSymbol, symbols, tradeSize, setTradeSize }: {
exchange: string; setExchange: (v: string) => void;
symbol: string; setSymbol: (v: string) => void;
symbols: string[];
tradeSize: number; setTradeSize: (v: number) => void;
}) {
const [timeframe, setTimeframe] = useState('30m');
const [days, setDays] = useState(7);
const [result, setResult] = useState<BacktestResult | null>(null);
const [loading, setLoading] = useState(false);
const [error, setError] = useState('');
const runBacktest = useCallback(async () => {
setLoading(true);
setError('');
@@ -110,12 +203,10 @@ export default function BacktestPage() {
} finally {
setLoading(false);
}
}, [symbol, timeframe, days, tradeSize]);
}, [symbol, exchange, timeframe, days, tradeSize]);
return (
<div className="mx-auto max-w-[1000px] bg-bg-primary p-5 text-text-primary">
<h1 className="text-2xl font-bold text-text-heading">📊 {t('Backtest')}</h1>
<div>
{/* Controls */}
<div className="mb-5 flex flex-wrap items-end gap-3">
<div>
@@ -135,7 +226,7 @@ export default function BacktestPage() {
<div>
<label className="text-xs text-text-secondary">Timeframe</label><br />
<select value={timeframe} onChange={e => setTimeframe(e.target.value)} className={selectClass}>
{TIMEFRAMES.map(t => <option key={t} value={t}>{t}</option>)}
{TIMEFRAMES.map(tf => <option key={tf} value={tf}>{tf}</option>)}
</select>
</div>
<div>
@@ -289,6 +380,266 @@ export default function BacktestPage() {
);
}
// ═══════════════ WALK-FORWARD ═══════════════
function EquityCurveSvg({ points }: { points: number[] }) {
if (points.length < 2) return <div className="text-xs text-text-dim">Not enough out-of-sample trades to chart.</div>;
const w = 600, h = 140, pad = 6;
const min = Math.min(...points), max = Math.max(...points);
const range = max - min || 1;
const stepX = (w - pad * 2) / (points.length - 1);
const toY = (v: number) => h - pad - ((v - min) / range) * (h - pad * 2);
const path = points.map((v, i) => `${i === 0 ? 'M' : 'L'} ${pad + i * stepX} ${toY(v)}`).join(' ');
const zeroY = toY(0);
const isPositive = points[points.length - 1] >= 0;
return (
<svg viewBox={`0 0 ${w} ${h}`} className="w-full" style={{ height: 140 }}>
<line x1={pad} y1={zeroY} x2={w - pad} y2={zeroY} stroke="var(--color-border-muted)" strokeDasharray="4 3" />
<path d={path} fill="none" stroke={isPositive ? 'var(--color-green)' : 'var(--color-red)'} strokeWidth={2} />
</svg>
);
}
function WalkForwardView({ exchange, setExchange, symbol, setSymbol, symbols, tradeSize, setTradeSize }: {
exchange: string; setExchange: (v: string) => void;
symbol: string; setSymbol: (v: string) => void;
symbols: string[];
tradeSize: number; setTradeSize: (v: number) => void;
}) {
const [timeframe, setTimeframe] = useState('4h');
const [totalDays, setTotalDays] = useState(1095);
const [trainDays, setTrainDays] = useState(270);
const [testDays, setTestDays] = useState(90);
const [showAdvanced, setShowAdvanced] = useState(false);
const [result, setResult] = useState<WalkForwardResult | null>(null);
const [loading, setLoading] = useState(false);
const [error, setError] = useState('');
const [needsLogin, setNeedsLogin] = useState(false);
const [history, setHistory] = useState<WfHistoryItem[]>([]);
const [historyLoading, setHistoryLoading] = useState(true);
const loadHistory = useCallback(async () => {
setHistoryLoading(true);
try {
const data = await apiFetch<WfHistoryItem[]>('/walk-forward/history?limit=20');
setHistory(data);
setNeedsLogin(false);
} catch (e: any) {
if (e instanceof ApiServiceError && e.status === 401) setNeedsLogin(true);
} finally {
setHistoryLoading(false);
}
}, []);
useEffect(() => { loadHistory(); }, [loadHistory]);
const runWalkForward = useCallback(async () => {
setLoading(true);
setError('');
setResult(null);
try {
const data = await apiFetch<WalkForwardResult>(
`/walk-forward/run?symbol=${encodeURIComponent(symbol)}&exchange=${exchange}&timeframe=${timeframe}` +
`&total_days=${totalDays}&train_days=${trainDays}&test_days=${testDays}&trade_size=${tradeSize}`,
{ method: 'POST' },
);
setResult(data);
setNeedsLogin(false);
loadHistory();
} catch (e: any) {
if (e instanceof ApiServiceError && e.status === 401) {
setNeedsLogin(true);
} else {
setError(e.message);
}
} finally {
setLoading(false);
}
}, [symbol, exchange, timeframe, totalDays, trainDays, testDays, tradeSize, loadHistory]);
const loadPastRun = useCallback(async (id: string) => {
setLoading(true);
setError('');
try {
const data = await apiFetch<WalkForwardResult>(`/walk-forward/${id}`);
setResult(data);
} catch (e: any) {
setError(e.message);
} finally {
setLoading(false);
}
}, []);
const deletePastRun = useCallback(async (id: string) => {
try {
await apiFetch(`/walk-forward/${id}`, { method: 'DELETE' });
loadHistory();
} catch {}
}, [loadHistory]);
return (
<div>
<div className="mb-4 rounded-lg border border-accent-blue/25 bg-accent-blue/10 p-3 text-xs text-text-secondary">
💡 Tối ưu tự động ngưỡng tín hiệu &amp; thời gian giữ lệnh trên từng cửa sổ dữ liệu quá khứ (train), rồi kiểm định trên dữ liệu chưa từng thấy (test) — tránh overfitting so với chạy 1 lần trên toàn bộ lịch sử.
</div>
{/* Controls */}
<div className="mb-3 flex flex-wrap items-end gap-3">
<div>
<label className="text-xs text-text-secondary">Exchange</label><br />
<select value={exchange} onChange={e => setExchange(e.target.value)} className={selectClass}>
{EXCHANGES.map(ex => <option key={ex} value={ex}>{ex.toUpperCase()}</option>)}
</select>
</div>
<div>
<label className="text-xs text-text-secondary">Symbol ({symbols.length})</label><br />
<select value={symbol} onChange={e => setSymbol(e.target.value)} className={`${selectClass} min-w-[140px]`}>
{symbols.length > 0
? symbols.map(s => <option key={s} value={s}>{s}</option>)
: <option value="">Loading...</option>}
</select>
</div>
<div>
<label className="text-xs text-text-secondary">Timeframe</label><br />
<select value={timeframe} onChange={e => setTimeframe(e.target.value)} className={selectClass}>
{WF_TIMEFRAMES.map(tf => <option key={tf} value={tf}>{tf}</option>)}
</select>
</div>
<div>
<label className="text-xs text-text-secondary">Trade Size (USDT)</label><br />
<input type="number" value={tradeSize} onChange={e => setTradeSize(Number(e.target.value))} className={`${selectClass} w-20`} min={1} />
</div>
<button
onClick={runWalkForward}
disabled={loading}
className={`rounded-md border-none bg-accent-blue px-6 py-2 font-semibold text-white ${loading ? 'cursor-wait' : 'cursor-pointer'}`}
>
{loading ? '⏳ Đang chạy...' : '🧪 Run Walk-Forward'}
</button>
</div>
<button onClick={() => setShowAdvanced(!showAdvanced)} className="mb-3 text-xs text-accent underline">
{showAdvanced ? '▾ Ẩn tùy chọn nâng cao' : '▸ Tùy chọn nâng cao (window size)'}
</button>
{showAdvanced && (
<div className={`${cardClass} mb-4 flex flex-wrap gap-4`}>
<div>
<label className="text-xs text-text-secondary">Tổng dữ liệu (ngày)</label><br />
<input type="number" value={totalDays} onChange={e => setTotalDays(Number(e.target.value))} className={`${selectClass} w-24`} min={180} />
</div>
<div>
<label className="text-xs text-text-secondary">Train window (ngày)</label><br />
<input type="number" value={trainDays} onChange={e => setTrainDays(Number(e.target.value))} className={`${selectClass} w-24`} min={30} />
</div>
<div>
<label className="text-xs text-text-secondary">Test window (ngày)</label><br />
<input type="number" value={testDays} onChange={e => setTestDays(Number(e.target.value))} className={`${selectClass} w-24`} min={14} />
</div>
</div>
)}
{needsLogin && <div className="mb-4 rounded-md border border-yellow/25 bg-yellow/10 p-3 text-sm text-yellow">🔒 Đăng nhập để chạy và lưu Walk-Forward Analysis.</div>}
{error && <div className="mb-4 text-red">❌ {error}</div>}
{result && (
<>
<div className="mb-5 grid grid-cols-[repeat(auto-fit,minmax(140px,1fr))] gap-3">
<SummaryCard label="OOS Trades" value={result.out_of_sample_summary.trades.toString()} />
<SummaryCard label="OOS Win Rate" value={`${result.out_of_sample_summary.win_rate}%`}
color={result.out_of_sample_summary.win_rate >= 50 ? 'text-green' : 'text-red'} />
<SummaryCard label="OOS Total PnL" value={`$${result.out_of_sample_summary.total_pnl.toFixed(2)}`}
color={result.out_of_sample_summary.total_pnl >= 0 ? 'text-green' : 'text-red'} />
<SummaryCard label="OOS Profit Factor" value={result.out_of_sample_summary.profit_factor?.toFixed(2) ?? '∞'}
color={result.out_of_sample_summary.profit_factor && result.out_of_sample_summary.profit_factor >= 1 ? 'text-green' : 'text-red'} />
<SummaryCard label="Max Drawdown" value={`${result.out_of_sample_summary.max_drawdown_pct}%`} color="text-red" />
<SummaryCard label="Folds" value={result.folds.length.toString()} />
</div>
<div className={`${cardClass} mb-4`}>
<h3 className="mb-3 text-text-heading">📈 Out-of-Sample Equity Curve (đã ghép các fold)</h3>
<EquityCurveSvg points={result.out_of_sample_summary.equity_curve} />
</div>
<div className={cardClass}>
<h3 className="mb-1 text-text-heading">📋 Chi tiết từng Fold</h3>
<p className="mb-3 text-xs text-text-dim">So sánh In-Sample (train, đã tối ưu) với Out-of-Sample (test, chưa từng thấy) — chênh lệch càng lớn thì càng có dấu hiệu overfitting.</p>
<div className="overflow-x-auto">
<table className="w-full min-w-[720px] border-collapse text-[12px]">
<thead>
<tr className="border-b border-border-default">
<th className={thClass}>Fold</th>
<th className={thClass}>Train</th>
<th className={thClass}>Test</th>
<th className={thClass}>Best Params</th>
<th className={thClass}>In-Sample WR / PnL</th>
<th className={thClass}>Out-of-Sample WR / PnL</th>
</tr>
</thead>
<tbody>
{result.folds.map(f => (
<tr key={f.fold_index} className="border-b border-border-muted">
<td className={tdClass}>#{f.fold_index + 1}</td>
<td className={tdClass}>{f.train_start.slice(0, 10)} → {f.train_end.slice(0, 10)}</td>
<td className={tdClass}>{f.test_start.slice(0, 10)} → {f.test_end.slice(0, 10)}</td>
<td className={`${tdClass} text-text-dim`}>
S≥{f.best_params.strong_threshold} / B≥{f.best_params.signal_threshold} / {f.best_params.max_hold_candles}c
</td>
<td className={tdClass}>{f.in_sample.win_rate}% / <span className={f.in_sample.total_pnl >= 0 ? 'text-green' : 'text-red'}>${f.in_sample.total_pnl.toFixed(2)}</span></td>
<td className={tdClass}>{f.out_of_sample.win_rate}% / <span className={f.out_of_sample.total_pnl >= 0 ? 'text-green' : 'text-red'}>${f.out_of_sample.total_pnl.toFixed(2)}</span></td>
</tr>
))}
</tbody>
</table>
</div>
</div>
</>
)}
{/* History */}
{!needsLogin && (
<div className={`${cardClass} mt-4`}>
<h3 className="mb-3 text-text-heading">🕓 Lịch sử Walk-Forward</h3>
{historyLoading ? (
<div className="text-xs text-text-dim">Đang tải...</div>
) : history.length === 0 ? (
<div className="text-xs text-text-dim">Chưa có lần chạy nào được lưu.</div>
) : (
<div className="overflow-x-auto">
<table className="w-full min-w-[600px] border-collapse text-[12px]">
<thead>
<tr className="border-b border-border-default">
<th className={thClass}>Symbol</th>
<th className={thClass}>TF</th>
<th className={thClass}>Folds</th>
<th className={thClass}>OOS Win Rate</th>
<th className={thClass}>OOS PnL</th>
<th className={thClass}>Ngày</th>
<th className={thClass}></th>
</tr>
</thead>
<tbody>
{history.map(h => (
<tr key={h.id} className="border-b border-border-muted">
<td className={tdClass}>{h.symbol} @ {h.exchange}</td>
<td className={tdClass}>{h.timeframe}</td>
<td className={tdClass}>{h.folds_count}</td>
<td className={tdClass}>{h.oos_win_rate ?? '-'}%</td>
<td className={`${tdClass} ${(h.oos_total_pnl ?? 0) >= 0 ? 'text-green' : 'text-red'}`}>${(h.oos_total_pnl ?? 0).toFixed(2)}</td>
<td className={`${tdClass} text-text-dim`}>{h.created_at.slice(0, 10)}</td>
<td className={tdClass}>
<button onClick={() => loadPastRun(h.id)} className="mr-2 text-accent underline">Xem</button>
<button onClick={() => deletePastRun(h.id)} className="text-red underline">Xóa</button>
</td>
</tr>
))}
</tbody>
</table>
</div>
)}
</div>
)}
</div>
);
}
function SummaryCard({ label, value, color }: { label: string; value: string; color?: string }) {
return (
<div className="rounded-lg border border-border-default bg-bg-surface p-3.5">
+120
View File
@@ -0,0 +1,120 @@
# Theo dõi đánh giá dự án Trading Portal — v7
> **Ngày đánh giá gốc:** 2026-07-03
> **Cập nhật v1-v6:** xem các file `theo_doi_trading-portal_v1.md`…`v6.md`
> **Cập nhật v7 (lần này):** 2026-07-04 — xử lý hạng mục cuối cùng còn lại **(m) Walk-Forward Backtest Optimization**, đóng toàn bộ danh sách nhược điểm ban đầu (a→r) trừ 2FA (cố ý hoãn vô thời hạn theo yêu cầu người dùng)
> **Người thực hiện:** Claude (Sonnet 5), theo yêu cầu của tien.a.le@accenture.com
> **Quy ước đặt tên:** Mỗi lần có thay đổi lớn → tạo bản mới `theo_doi_trading-portal_v8.md`, ... giữ nguyên các bản cũ làm lịch sử.
---
## 1. Tổng quan dự án
(Không đổi — xem [theo_doi_trading-portal_v2.md](theo_doi_trading-portal_v2.md) mục 1.)
---
## 2. Walk-Forward Backtest Optimization (m) — thiết kế & triển khai
### 2.1 Vấn đề gốc
Hệ thống dùng backtest trên dữ liệu lịch sử để đánh giá chiến lược, nhưng các ngưỡng/tham số của 13 thuật toán vote trong `signal_scoring.py` lại được điều chỉnh thủ công dựa trên quan sát chính kết quả backtest đó — tạo vòng lặp look-ahead bias: dữ liệu dùng để đánh giá cũng là dữ liệu dùng để tinh chỉnh, khiến kết quả backtest "đẹp" hơn thực tế sẽ chạy live.
### 2.2 Quyết định thiết kế (đã thống nhất với người dùng)
| Câu hỏi | Quyết định |
|---|---|
| Dữ liệu bao lâu? | Giữ 3 năm (`total_days=1095` mặc định) — đủ chia ~9 fold train/test có ý nghĩa thống kê ở timeframe 4h/1h. Token mới niêm yết tự dùng bao nhiêu dữ liệu có sẵn thay vì báo lỗi. |
| Tự động tối ưu tham số? | Có — grid search trên 3 tham số rẻ để thử: `strong_threshold`, `signal_threshold` (ngưỡng phân loại tín hiệu STRONG/thường), `max_hold_candles` (thời gian giữ lệnh tối đa). **Không** đụng vào các hằng số nội bộ của 13 thuật toán (RSI 70/30, MFI 20/80, v.v.) — refactor toàn bộ số đó rủi ro cao hơn lợi ích, để lại cho một đợt riêng nếu cần. |
| Win-rate cache của `signal_booster.py` — chuyển sang out-of-sample? | **Giữ nguyên, không đổi.** Cache đó đã tính từ `hypothetical_trades` chạy live thật (không phải dữ liệu backtest) — đã có suy giảm hàm mũ (half-life 14 ngày, tự thích nghi regime), đã có ngưỡng mẫu tối thiểu (~15 trade). Đây thực chất **đã là out-of-sample thật**, đáng tin hơn cả out-of-sample của walk-forward (vốn vẫn là dữ liệu lịch sử). Walk-forward dùng để kiểm định/tối ưu tham số thiết kế chiến lược; cache này dùng để thích nghi liên tục khi chạy thật — bổ trợ nhau, không thay thế. |
| UI/report? | Thêm tab "Walk-Forward" trong trang `/backtest` (không đụng bản Backtest đơn giản trong Profile) — bảng kết quả từng fold (train/test, tham số tốt nhất, so sánh in-sample vs out-of-sample), đường equity curve out-of-sample ghép từ tất cả fold, và lưu lịch sử các lần chạy. |
### 2.3 Cách hoạt động
1. Chia 3 năm dữ liệu thành các **fold trượt**: cửa sổ train cố định (mặc định 270 ngày) trượt tới theo bước = độ dài cửa sổ test (mặc định 90 ngày) → ~9 fold.
2. Với mỗi fold: **grid search** 27 tổ hợp tham số (3×3×3) trên cửa sổ **train**, chọn tổ hợp tốt nhất theo hàm mục tiêu kiểu Sharpe (trung bình PnL/độ lệch chuẩn, nhân √n để phạt số lệnh quá ít) — không chọn theo tổng PnL thô để tránh bị 1 lệnh may mắn chi phối.
3. Áp tham số tốt nhất đó vào cửa sổ **test** (dữ liệu chưa từng dùng để tối ưu) → kết quả out-of-sample của fold.
4. Ghép toàn bộ kết quả out-of-sample của các fold theo thời gian → equity curve, win rate, profit factor, max drawdown tổng — đây là con số đáng tin nhất, vì nó chưa từng "nhìn thấy" dữ liệu nó được đánh giá trên.
### 2.4 Tối ưu hiệu năng quan trọng
13 thuật toán vote (BB/RSI, MACD, SuperTrend, MFI, v.v.) không phụ thuộc vào `strong_threshold`/`signal_threshold` — chỉ bước phân loại cuối cùng mới phụ thuộc. Đã tách `_classify_signal_combined` trong `signal_scoring.py` thành `_compute_adjusted_score` (phần đắt, tính 1 lần) + `_score_to_signal` (phần rẻ, chỉ so sánh ngưỡng). Nhờ vậy grid search 27 tổ hợp tham số trên 1 fold chỉ tốn thêm chi phí "replay ngưỡng" rất rẻ, không phải chạy lại toàn bộ 13 thuật toán 27 lần.
### 2.5 Các file mới/thay đổi
- `backend/app/services/signal_scoring.py` — tách `_compute_adjusted_score` + `_score_to_signal` khỏi `_classify_signal_combined` (tương thích ngược 100%, live trading không đổi hành vi).
- `backend/app/services/backtest_engine.py` (MỚI) — chuyển toàn bộ engine backtest (fetch candle, precompute indicator, simulate trade) từ `api/v1/backtest.py` sang service layer đúng vị trí kiến trúc (tránh việc `walk_forward.py` phải import ngược từ tầng API — cùng tinh thần tách "god file" như mục (h) ở v6).
- `backend/app/api/v1/backtest.py` — giờ chỉ còn route handler, gọi vào `backtest_engine`.
- `backend/app/services/walk_forward.py` (MỚI) — fold generation, grid search, out-of-sample aggregation.
- `backend/app/api/v1/walk_forward.py` (MỚI) — `POST /walk-forward/run` (tự lưu), `GET /walk-forward/history`, `GET /walk-forward/{id}`, `DELETE /walk-forward/{id}`. Yêu cầu đăng nhập.
- `backend/alembic/versions/merge_heads_1.py` (MỚI) — **phát hiện phụ**: lịch sử Alembic đã bị phân nhánh thành 2 head không hợp nhất (`1b3f1630986f` và `4_add_sl_tp_columns`, cả hai đều rẽ từ `add_candle_partitions`) từ trước, khiến `alembic upgrade head` sẽ lỗi mơ hồ trên môi trường mới. Đã tạo migration merge (no-op) để gộp lại thành 1 head duy nhất trước khi thêm bảng mới.
- `backend/alembic/versions/add_walk_forward_results.py` (MỚI) — bảng `walk_forward_results` lưu lịch sử các lần chạy walk-forward.
- `frontend/src/features/backtest/BacktestPage.tsx` — thêm toggle "Single Run / Walk-Forward"; tab Walk-Forward có control (Exchange/Symbol/Timeframe giới hạn 1h-4h để giữ thời gian chạy nhanh/Trade Size), tùy chọn nâng cao (window size), bảng chi tiết từng fold, equity curve SVG, và lịch sử các lần chạy.
- **Test mới:** `test_backtest_engine.py` (8 test) + `test_walk_forward.py` (11 test) — tổng **153 test pass** (tăng từ 134).
### 2.6 Giới hạn đã biết (cố ý, có ghi chú trong code)
- Chỉ tối ưu 3 tham số ngưỡng/thời gian giữ lệnh, không tối ưu các hằng số nội bộ của 13 thuật toán — xem mục 2.2.
- Timeframe walk-forward giới hạn 1h/4h trên UI (không cho 15m/30m) để giữ runtime nhanh — thuật toán backtest gốc dùng slicing mảng O(n²) mỗi candle, ở 4h/1h trên 3 năm vẫn đủ nhanh cho 1 request đồng bộ, nhưng ở 15m/30m số nến tăng gấp 4-8 lần sẽ chậm đáng kể. Không sửa thuật toán slicing gốc (out of scope, rủi ro cao hơn lợi ích ở đây).
- Chưa build async job/polling — `/walk-forward/run` chạy đồng bộ. Với default 3 năm/4h/9 fold thì đủ nhanh; nếu người dùng chỉnh nâng cao để chạy timeframe 1h với total_days rất lớn có thể chậm hơn — chưa có giới hạn cứng, chỉ giới hạn `total_days` trong khoảng [180, 1825] qua validation.
---
## 3. Toàn bộ nhược điểm & rủi ro — trạng thái tổng hợp đến v7
| # | Vấn đề | Trạng thái |
|---|---|---|
| a | Mật khẩu Gitea lộ trong lịch sử git | ✅ Đóng hoàn toàn (v5) |
| b | RBAC thiếu ở `/orders/place` | ✅ Đã xử lý + có test |
| c | Hardcode sàn "mexc" trong order routing | ✅ Đã xử lý + có test |
| d | AES-CBC không xác thực toàn vẹn | ✅ Đã chuyển sang AES-GCM + có test |
| e | Thiếu test suite/CI | 🟢 153 test, CI workflow đã thêm (chưa xác nhận runner) |
| f | Mật khẩu DB không đồng bộ trong docker-compose | ✅ Đã xử lý |
| g | CORS fallback `*` | ✅ Đã fail-closed + có test |
| h | God files (`signal_service.py`, `backtest.py`) | ✅ Đã tách `signal_scoring.py` (v6) + `backtest_engine.py` (v7) |
| i | Frontend thiếu tầng data-fetching thống nhất | ✅ Đã hợp nhất về `apiFetch` |
| j | AnalyticsPage dùng data giả | ✅ Đã nối vào `/analytics/dashboard` thật |
| k | Inline CSS-in-JS không design system | ✅ Đã có Tailwind design system, 14 file chuyển đổi |
| l | Cache in-memory single-instance | ✅ Đã thêm Redis + fallback graceful |
| m | Rủi ro overfitting hệ thống tín hiệu | ✅ **Walk-Forward Backtest Optimization đã triển khai (mục 2)** |
| n | Không có backup/restore Postgres | ✅ Đã có script backup/restore |
| o | Quản lý secrets không nhất quán | ✅ Đã chuyển sang Docker secrets pattern |
| p | Eviction dùng nhầm giá cross-symbol | ✅ Đã xử lý + có test |
| q | RSI sai giá trị khi giá đi ngang | ✅ Đã sửa + có test |
| r | MFI wraparound index | ✅ Đã sửa + có test |
| s | Sự cố quy trình: replace_all bỏ sót 1 vị trí | ✅ Đã vá, rút kinh nghiệm |
| t | Bảng thiếu overflow-x-auto trên mobile | ✅ Đã sửa (v6) |
| u | Grid 2 cột không responsive trong ProfilePage | ✅ Đã sửa (v6) |
| v | ProfilePage thiếu nav bar/logout | ✅ Đã sửa (v6) |
| w | Link `<a href>` thường gây full reload ở AdminPage | ⏳ Ghi nhận, không ưu tiên |
| x | refreshAccessToken() ép logout khi backend lỗi tạm thời | ⏳ Ghi nhận, cần bàn thiết kế riêng |
| y | (mới) Alembic có 2 head phân nhánh không hợp nhất | ✅ Đã merge (mục 2.5) |
**Toàn bộ danh sách gốc (a→r) từ bản đánh giá v0 nay đã được xử lý.** Việc còn lại ngoài danh sách gốc: bật 2FA cho tài khoản Gitea (cố ý hoãn vô thời hạn theo yêu cầu "chưa cần thiết cho hiện tại"), và 2 phát hiện phụ (w, x) không ưu tiên.
---
## 4. Đề xuất tiếp theo (không còn mục nào cấp thiết)
| Ưu tiên | Việc cần làm |
|---|---|
| Khi cần | Chạy thử Walk-Forward trên vài symbol thật để xem tham số tối ưu có ổn định qua các fold không (nếu nhảy lung tung giữa các fold → dấu hiệu chiến lược không robust) |
| Khi cần | Bật 2FA cho các tài khoản ghi trên Gitea (không cấp thiết theo yêu cầu người dùng) |
| Khi cần | Bàn chiến lược retry cho refresh-token khi backend lỗi tạm thời (x) |
| Khi cần | Đổi link `<a href="/profile">` trong AdminPage sang điều hướng SPA (w) |
| Dài hạn | Nếu muốn tối ưu sâu hơn walk-forward: mở rộng tối ưu sang các hằng số nội bộ 13 thuật toán — cần refactor lớn `signal_scoring.py`, nên bàn riêng trước khi làm |
---
## 5. Lịch sử phiên bản
| Phiên bản | Ngày | Thay đổi |
|---|---|---|
| v0 | 2026-07-03 | Đánh giá tổng thể lần đầu |
| v1 | 2026-07-03 | Fix (b), (c), (f) |
| v2 | 2026-07-03 | Test cho (b)/(c), fix (d) AES-GCM, fix (g) CORS, hướng dẫn (a) |
| v3 | 2026-07-03 | Test risk_manager/trade_executor/signal_service (81 test), CI Gitea Actions, phát hiện (p) |
| v4 | 2026-07-03 | Rewrite lịch sử git (a) chuẩn bị xong, fix (p), test indicator_service + async signal_service (122 test), phát hiện (q)/(r), ghi nhận sự cố quy trình (s) |
| v5 | 2026-07-03 | Đóng hoàn toàn sự cố (a) — rotate xong 4/4 mật khẩu, force-push lịch sử đã rewrite thành công, verify sạch |
| v6 | 2026-07-04 | Xử lý (h), (i), (j), (k), (l), (n), (o), (q), (r) — 134 test pass; Tailwind design system 14 file; thêm Redis; review UI, phát hiện & sửa (t, u, v), ghi nhận (w, x) |
| v7 | 2026-07-04 | **(m) Walk-Forward Backtest Optimization** — grid search tự động 3 tham số, out-of-sample stitching, UI tab mới trong `/backtest`; tách `backtest_engine.py` khỏi API layer (tiếp nối tinh thần (h)); phát hiện & sửa Alembic 2-head phân nhánh (y); 153 test pass. **Toàn bộ danh sách nhược điểm gốc a→r đã đóng.** |