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:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"}
|
||||
Reference in New Issue
Block a user