javitechjkd/backtestingv2
0
1"""
2API Routes for Backtesting Application
3"""
4
5from fastapi import APIRouter, HTTPException, WebSocket, WebSocketDisconnect
6from fastapi.responses import JSONResponse
7from typing import List, Optional
8from datetime import datetime
9import asyncio
10import logging
11import time
12import traceback
13import sys
14from pathlib import Path
15
16# Add parent directory for imports
17sys.path.insert(0, str(Path(__file__).parent.parent))
18
19from backtesting import Backtest
20import pandas as pd
21import numpy as np
22
23from api.schemas import (
24 BacktestConfig, BacktestRequest, BacktestResult, BacktestMetrics,
25 TradeResult, HealthResponse, StatusResponse, SymbolInfo, ConfigUpdate
26)
27from api.config_manager import get_config_manager
28from core.mt5_data_provider import get_data_provider
29from core.donchian_strategy import create_strategy
30from utils.validators import InputValidator
31from utils.timeout import async_timeout
32from utils.exceptions import ValidationException, TimeoutException, BacktestingException
33
34logger = logging.getLogger(__name__)
35
36router = APIRouter()
37
38# Global state for backtest status
39_backtest_state = {
40 "running": False,
41 "progress": 0.0,
42 "message": "Idle"
43}
44
45# WebSocket connections
46_ws_connections: List[WebSocket] = []
47
48
49def _safe_float(val) -> float:
50 """Ensure a value is a valid JSON-compliant float (no NaN or Inf)"""
51 try:
52 f_val = float(val)
53 if np.isnan(f_val) or np.isinf(f_val):
54 return 0.0
55 return f_val
56 except (ValueError, TypeError):
57 return 0.0
58
59
60@router.get("/health", response_model=HealthResponse)
61async def health_check():
62 """Health check endpoint"""
63 provider = get_data_provider()
64 mt5_connected = provider.connect()
65
66 return HealthResponse(
67 status="healthy" if mt5_connected else "degraded",
68 mt5_connected=mt5_connected,
69 timestamp=datetime.now()
70 )
71
72
73@router.get("/symbols")
74async def get_symbols():
75 """Get list of available trading symbols"""
76 provider = get_data_provider()
77 import MetaTrader5 as mt5
78
79 # Ensure connected
80 if not provider.connect():
81 return {"symbols": ["BTCUSD"]}
82
83 # Get all symbols from server
84 all_symbols_raw = mt5.symbols_get()
85 if all_symbols_raw is None or len(all_symbols_raw) == 0:
86 return {"symbols": ["BTCUSD"]}
87
88 # Get names of all available symbols
89 all_names = {s.name for s in all_symbols_raw}
90
91 # Get names of visible symbols (Market Watch)
92 visible_names = [s.name for s in all_symbols_raw if s.visible]
93
94 # Top priority: Common symbols requested by user
95 common_symbols = [
96 "BTCUSD", "ETHUSD", "SOLUSD",
97 "TSLA", "ADBE", "GOOGL", "AAPL", "MSFT", "NVDA",
98 "EURUSD", "GBPUSD", "USDJPY"
99 ]
100
101 # Build final list: Common first, then visible, then rest (up to 50)
102 final_list = []
103 seen = set()
104
105 # 1. Add common symbols if they exist on server
106 for s in common_symbols:
107 if s in all_names and s not in seen:
108 final_list.append(s)
109 seen.add(s)
110
111 # 2. Add visible symbols (Market Watch favorites)
112 for s in visible_names:
113 if s not in seen:
114 final_list.append(s)
115 seen.add(s)
116
117 # 3. Add some other available symbols to fill up to a reasonable number
118 for s_raw in all_symbols_raw[:100]:
119 if s_raw.name not in seen:
120 final_list.append(s_raw.name)
121 seen.add(s_raw.name)
122 if len(final_list) >= 50:
123 break
124
125 return {"symbols": final_list}
126
127
128@router.get("/symbol/{symbol}", response_model=SymbolInfo)
129async def get_symbol_info(symbol: str):
130 """Get information about a specific symbol"""
131 provider = get_data_provider()
132 info = provider.get_symbol_info(symbol)
133
134 if info is None:
135 raise HTTPException(status_code=404, detail=f"Symbol {symbol} not found")
136
137 return SymbolInfo(**info)
138
139
140@router.get("/timeframes")
141async def get_timeframes():
142 """Get available timeframes"""
143 return {
144 "timeframes": [
145 {"value": "D1", "label": "Daily", "description": "1 Day"},
146 {"value": "W1", "label": "Weekly", "description": "1 Week"},
147 {"value": "MN", "label": "Monthly", "description": "1 Month"},
148 {"value": "H4", "label": "4 Hours", "description": "4 Hours"},
149 {"value": "H1", "label": "1 Hour", "description": "1 Hour"},
150 ]
151 }
152
153
154@router.get("/config", response_model=BacktestConfig)
155async def get_config():
156 """Get current configuration"""
157 manager = get_config_manager()
158 return manager.get_config()
159
160
161@router.put("/config", response_model=BacktestConfig)
162async def update_config(update: ConfigUpdate):
163 """Update configuration (hot-reload)"""
164 manager = get_config_manager()
165 manager.update_config(update.config)
166
167 # Notify WebSocket clients
168 await broadcast_message({
169 "type": "config_updated",
170 "config": update.config.model_dump()
171 })
172
173 return manager.get_config()
174
175
176@router.patch("/config")
177async def patch_config(updates: dict):
178 """Partially update configuration"""
179 manager = get_config_manager()
180 manager.update_partial(updates)
181
182 # Notify WebSocket clients
183 await broadcast_message({
184 "type": "config_updated",
185 "config": manager.get_config().model_dump()
186 })
187
188 return manager.get_config()
189
190
191@router.post("/config/reset")
192async def reset_config():
193 """Reset configuration to defaults"""
194 manager = get_config_manager()
195 manager.reset_to_defaults()
196 return manager.get_config()
197
198
199@router.get("/backtest/status", response_model=StatusResponse)
200async def get_backtest_status():
201 """Get current backtest status"""
202 return StatusResponse(**_backtest_state)
203
204
205@router.post("/backtest/run", response_model=BacktestResult)
206@async_timeout(300) # 5 minute timeout for backtest
207async def run_backtest(request: BacktestRequest):
208 """Run a backtest with the given configuration"""
209 global _backtest_state
210
211 if _backtest_state["running"]:
212 raise HTTPException(status_code=409, detail="Backtest already running")
213
214 config = request.config
215 start_time = time.time()
216
217 try:
218 # VALIDATION: Validate backtest parameters
219 try:
220 InputValidator.validate_backtest_params(
221 initial_capital=config.initial_capital,
222 risk_percent=config.risk_percent,
223 symbol=config.symbol,
224 bars=config.bars
225 )
226 except ValidationException as e:
227 logger.warning(f"Validation error: {e.message}")
228 raise HTTPException(
229 status_code=400,
230 detail={
231 "error": "VALIDATION_ERROR",
232 "message": e.message,
233 "details": e.details
234 }
235 )
236
237 _backtest_state = {"running": True, "progress": 0.0, "message": "Initializing..."}
238 await broadcast_message({"type": "status", **_backtest_state})
239
240 # Get data from MT5
241 _backtest_state["message"] = f"Fetching data for {config.symbol}..."
242 _backtest_state["progress"] = 10.0
243 await broadcast_message({"type": "status", **_backtest_state})
244
245 provider = get_data_provider()
246 # Convert timeframe enum to string value
247 timeframe_str = config.timeframe.value if hasattr(config.timeframe, 'value') else str(config.timeframe)
248 data = provider.get_data(
249 symbol=config.symbol,
250 timeframe=timeframe_str,
251 bars=config.bars
252 )
253
254 if data is None or data.empty:
255 raise HTTPException(status_code=500, detail=f"Failed to fetch data for {config.symbol}")
256
257 _backtest_state["message"] = f"Running backtest on {len(data)} bars..."
258 _backtest_state["progress"] = 30.0
259 await broadcast_message({"type": "status", **_backtest_state})
260
261 # Create strategy with config
262 # Extract enum values properly
263 trade_dir = config.trade_direction.value if hasattr(config.trade_direction, 'value') else str(config.trade_direction)
264 trail_mode = config.trailing_mode.value if hasattr(config.trailing_mode, 'value') else str(config.trailing_mode)
265
266 first_price = data['Open'].iloc[0] if not data.empty else 1.0
267 price_scaling_factor = first_price / 100.0 if first_price > 0 else 1.0
268 # Ensure we don't have a zero scaling factor
269 if price_scaling_factor == 0: price_scaling_factor = 1.0
270
271 strategy_config = {
272 'donchian_period': config.donchian_period,
273 'risk_percent': config.risk_percent,
274 'break_margin': config.get_break_margin(),
275 'trade_direction': trade_dir,
276 'lookbehind': config.lookbehind,
277 'stop_resets_support': config.stop_resets_support,
278 'trailing_mode': trail_mode,
279 'leverage': config.leverage,
280 'offset': config.offset / price_scaling_factor,
281 'slippage': config.slippage / price_scaling_factor,
282 }
283
284 StrategyClass = create_strategy(strategy_config)
285
286 # Scale price down to allow fractional unit precision in backtesting.py
287 scaled_data = data.copy()
288 for col in ['Open', 'High', 'Low', 'Close']:
289 scaled_data[col] = data[col] / price_scaling_factor
290
291 # Calculate total commission (base commission % + slippage convert to %)
292 # Slippage in points (e.g. 2 points = $2 for BTC)
293 slippage_percent = (config.slippage / first_price) if first_price > 0 else 0
294 total_commission = (config.commission / 100.0) + slippage_percent
295
296 # Create Backtest instance with scaled data
297 bt = Backtest(
298 scaled_data,
299 StrategyClass,
300 cash=config.initial_capital,
301 commission=total_commission,
302 margin=1 / config.leverage,
303 exclusive_orders=True
304 )
305
306 _backtest_state["message"] = "Executing strategy..."
307 _backtest_state["progress"] = 50.0
308 await broadcast_message({"type": "status", **_backtest_state})
309
310 stats = bt.run()
311
312 # Scale back prices in results and indicators
313 trades_df = stats.get('_trades')
314 if trades_df is not None and len(trades_df) > 0:
315 trades_df['EntryPrice'] *= price_scaling_factor
316 trades_df['ExitPrice'] *= price_scaling_factor
317
318 # Scale back indicators in strategy for correct visualization
319 strat = stats['_strategy']
320 if hasattr(strat, 'donchian_high'): strat.donchian_high[:] *= price_scaling_factor
321 if hasattr(strat, 'donchian_low'): strat.donchian_low[:] *= price_scaling_factor
322 if hasattr(strat, 'res_indicator'): strat.res_indicator[:] *= price_scaling_factor
323 if hasattr(strat, 'sup_indicator'): strat.sup_indicator[:] *= price_scaling_factor
324 if hasattr(strat, 'sl_indicator'): strat.sl_indicator[:] *= price_scaling_factor
325
326 # Processing results
327 final_equity = float(stats.get('Equity Final [$]', config.initial_capital))
328 total_return = final_equity - config.initial_capital
329 total_return_percent = (total_return / config.initial_capital) * 100 if config.initial_capital > 0 else 0
330
331 # LOGGING FOR DEBUGGING
332 with open('backtest_debug.log', 'a') as f:
333 f.write(f"\n--- Backtest Result {datetime.now()} ---\n")
334 f.write(f"Symbol: {config.symbol}, Timeframe: {config.timeframe}\n")
335 f.write(f"Initial Capital: {config.initial_capital}\n")
336 f.write(f"Final Equity: {final_equity}\n")
337 f.write(f"Total Return: {total_return}\n")
338 f.write(f"Return %: {total_return_percent}\n")
339 f.write(f"Sharpe: {stats.get('Sharpe Ratio')}\n")
340 f.write("-" * 40 + "\n")
341
342 # Extract metrics
343 metrics = BacktestMetrics(
344 total_return=_safe_float(total_return),
345 total_return_percent=_safe_float(total_return_percent),
346 sharpe_ratio=_safe_float(stats.get('Sharpe Ratio', 0)),
347 max_drawdown=_safe_float(stats.get('Max. Drawdown [$]', 0)) if 'Max. Drawdown [$]' in stats else (_safe_float(stats.get('Max. Drawdown [%]', 0)) * config.initial_capital / 100),
348 max_drawdown_percent=_safe_float(stats.get('Max. Drawdown [%]', 0)),
349 win_rate=_safe_float(stats.get('Win Rate [%]', 0)),
350 profit_factor=_safe_float(stats.get('Profit Factor', 0)),
351 total_trades=int(stats.get('# Trades', 0)),
352 winning_trades=int(stats.get('# Trades', 0) * stats.get('Win Rate [%]', 0) / 100) if stats.get('# Trades', 0) > 0 else 0,
353 losing_trades=int(stats.get('# Trades', 0) - (stats.get('# Trades', 0) * stats.get('Win Rate [%]', 0) / 100)) if stats.get('# Trades', 0) > 0 else 0,
354 avg_trade_return=_safe_float(stats.get('Avg. Trade [%]', 0)),
355 avg_trade_duration=str(stats.get('Avg. Trade Duration', 'N/A')),
356 avg_winning_trade=_safe_float(stats.get('Best Trade [%]', 0)) if stats.get('# Trades', 0) > 0 else 0,
357 avg_losing_trade=_safe_float(stats.get('Worst Trade [%]', 0)) if stats.get('# Trades', 0) > 0 else 0,
358 best_trade=_safe_float(stats.get('Best Trade [%]', 0)),
359 worst_trade=_safe_float(stats.get('Worst Trade [%]', 0)),
360 start_equity=_safe_float(config.initial_capital),
361 final_equity=_safe_float(final_equity),
362 exposure_time=_safe_float(stats.get('Exposure Time [%]', 0))
363 )
364
365 # Extract trades (closed)
366 trades_list = []
367 if '_trades' in stats:
368 trades_df = stats['_trades']
369 for _, trade in trades_df.iterrows():
370 trades_list.append(TradeResult(
371 entry_time=trade['EntryTime'],
372 exit_time=trade['ExitTime'],
373 entry_price=_safe_float(trade['EntryPrice']),
374 exit_price=_safe_float(trade['ExitPrice']),
375 size=_safe_float(trade['Size']) / price_scaling_factor,
376 pnl=_safe_float(trade['PnL']),
377 pnl_percent=_safe_float(trade['ReturnPct']) * 100 if 'ReturnPct' in trade else 0,
378 type="LONG" if trade['Size'] > 0 else "SHORT",
379 duration=str(trade['ExitTime'] - trade['EntryTime']),
380 return_percent=_safe_float(trade['ReturnPct']) * 100 if 'ReturnPct' in trade else 0,
381 is_open=False
382 ))
383
384 # Extract open positions (if any)
385 strategy = stats.get('_strategy')
386 if strategy and hasattr(strategy, 'trades'):
387 for active_trade in strategy.trades:
388 # Only include trades that are still open
389 if hasattr(active_trade, 'size') and active_trade.size != 0:
390 entry_price = _safe_float(active_trade.entry_price) * price_scaling_factor
391 current_price = _safe_float(data['Close'].iloc[-1])
392 size = _safe_float(active_trade.size) / price_scaling_factor
393
394 # Calculate unrealized P&L
395 if active_trade.is_long:
396 unrealized_pnl = (current_price - entry_price) * abs(size)
397 else:
398 unrealized_pnl = (entry_price - current_price) * abs(size)
399
400 pnl_percent = (unrealized_pnl / config.initial_capital) * 100 if config.initial_capital > 0 else 0
401
402 # Get current SL
403 current_sl = _safe_float(active_trade.sl) * price_scaling_factor if active_trade.sl else None
404
405 trades_list.append(TradeResult(
406 entry_time=active_trade.entry_time if hasattr(active_trade, 'entry_time') else data.index[0],
407 exit_time=None, # Open position
408 entry_price=entry_price,
409 exit_price=None, # Open position
410 size=size,
411 pnl=_safe_float(unrealized_pnl),
412 pnl_percent=_safe_float(pnl_percent),
413 type="LONG" if active_trade.is_long else "SHORT",
414 duration=None,
415 return_percent=_safe_float(pnl_percent),
416 is_open=True,
417 current_sl=current_sl
418 ))
419
420 # Extract equity curve
421 equity_curve = []
422 if '_equity_curve' in stats:
423 eq = stats['_equity_curve']
424 for idx, row in eq.iterrows():
425 # Use strftime to avoid Windows Errno 22 with isoformat()
426 try:
427 time_str = idx.strftime('%Y-%m-%dT%H:%M:%S') if hasattr(idx, 'strftime') else str(idx)
428 except (OSError, ValueError):
429 time_str = str(idx)
430
431 equity_curve.append({
432 "time": time_str,
433 "equity": _safe_float(row['Equity']),
434 "drawdown": _safe_float(row.get('DrawdownPct', 0)) * 100 if 'DrawdownPct' in row else 0
435 })
436
437 # Prepare chart data
438 chart_data = _prepare_chart_data(data, stats, config)
439
440 execution_time = time.time() - start_time
441
442 _backtest_state = {"running": False, "progress": 100.0, "message": "Complete"}
443 await broadcast_message({"type": "status", **_backtest_state})
444
445 result = BacktestResult(
446 success=True,
447 config=config,
448 metrics=metrics,
449 trades=trades_list,
450 equity_curve=equity_curve,
451 chart_data=chart_data,
452 execution_time=execution_time
453 )
454
455 await broadcast_message({
456 "type": "result",
457 "data": result.model_dump()
458 })
459
460 return result
461
462 except ValidationException as e:
463 # Handle validation errors gracefully
464 logger.warning(f"Validation error in backtest: {e.message}")
465 _backtest_state = {"running": False, "progress": 0.0, "message": f"Validation error: {e.message}"}
466 await broadcast_message({"type": "status", **_backtest_state})
467
468 raise HTTPException(
469 status_code=400,
470 detail={
471 "error": e.error_code,
472 "message": e.message,
473 "details": e.details
474 }
475 )
476 except TimeoutException as e:
477 # Handle timeout errors
478 logger.error(f"Timeout in backtest: {e.message}")
479 _backtest_state = {"running": False, "progress": 0.0, "message": f"Timeout: {e.message}"}
480 await broadcast_message({"type": "status", **_backtest_state})
481
482 raise HTTPException(
483 status_code=504,
484 detail={
485 "error": e.error_code,
486 "message": e.message,
487 "details": e.details
488 }
489 )
490 except HTTPException:
491 raise
492 except Exception as e:
493 import sys
494 print(f"CRITICAL BACKTEST ERROR: {e}", file=sys.stderr)
495 full_traceback = traceback.format_exc()
496 print(full_traceback, file=sys.stderr)
497
498 logger.error(f"Backtest error: {e}")
499 logger.error(full_traceback)
500
501 _backtest_state = {"running": False, "progress": 0.0, "message": f"Error: {str(e)}"}
502 await broadcast_message({"type": "status", **_backtest_state})
503
504 return BacktestResult(
505 success=False,
506 config=config,
507 error=f"{str(e)}\n{full_traceback}",
508 execution_time=time.time() - start_time
509 )
510
511
512def _prepare_chart_data(data: pd.DataFrame, stats, config: BacktestConfig) -> dict:
513 """Prepare chart data for frontend visualization"""
514 strategy = stats['_strategy']
515
516 # Extract indicators from strategy
517 d_high = strategy.donchian_high
518 d_low = strategy.donchian_low
519 s_res = strategy.res_indicator
520 s_sup = strategy.sup_indicator
521 sl_line = strategy.sl_indicator
522 tendency = strategy.tend_indicator
523
524 # OHLC data
525 ohlc = []
526 donchian = []
527 structural = []
528 stop_loss = []
529
530 for i, (idx, row) in enumerate(data.iterrows()):
531 # Use strftime to avoid Windows Errno 22 with isoformat()
532 try:
533 time_str = idx.strftime('%Y-%m-%dT%H:%M:%S') if hasattr(idx, 'strftime') else str(idx)
534 except (OSError, ValueError):
535 time_str = str(idx)
536
537 ohlc.append({
538 "time": time_str,
539 "open": _safe_float(row['Open']),
540 "high": _safe_float(row['High']),
541 "low": _safe_float(row['Low']),
542 "close": _safe_float(row['Close']),
543 "volume": _safe_float(row.get('Volume', 0))
544 })
545
546 donchian.append({
547 "time": time_str,
548 "high": _safe_float(d_high[i]) if not np.isnan(d_high[i]) else None,
549 "low": _safe_float(d_low[i]) if not np.isnan(d_low[i]) else None
550 })
551
552 structural.append({
553 "time": time_str,
554 "resistance": _safe_float(s_res[i]) if not np.isnan(s_res[i]) else None,
555 "support": _safe_float(s_sup[i]) if not np.isnan(s_sup[i]) else None,
556 "tendency": int(tendency[i]) if not np.isnan(tendency[i]) else 0
557 })
558
559 stop_loss.append({
560 "time": time_str,
561 "sl": _safe_float(sl_line[i]) if not np.isnan(sl_line[i]) else None
562 })
563
564 # Trade signals
565 signals = []
566 if '_trades' in stats:
567 trades_df = stats['_trades']
568 for _, trade in trades_df.iterrows():
569 # Use strftime to avoid Windows Errno 22 with isoformat()
570 try:
571 entry_time_str = trade['EntryTime'].strftime('%Y-%m-%dT%H:%M:%S') if hasattr(trade['EntryTime'], 'strftime') else str(trade['EntryTime'])
572 except (OSError, ValueError):
573 entry_time_str = str(trade['EntryTime'])
574
575 try:
576 exit_time_str = trade['ExitTime'].strftime('%Y-%m-%dT%H:%M:%S') if hasattr(trade['ExitTime'], 'strftime') else str(trade['ExitTime'])
577 except (OSError, ValueError):
578 exit_time_str = str(trade['ExitTime'])
579
580 signals.append({
581 "time": entry_time_str,
582 "type": "entry",
583 "direction": "long" if trade['Size'] > 0 else "short",
584 "price": _safe_float(trade['EntryPrice'])
585 })
586 signals.append({
587 "time": exit_time_str,
588 "type": "exit",
589 "direction": "long" if trade['Size'] > 0 else "short",
590 "price": _safe_float(trade['ExitPrice'])
591 })
592
593 return {
594 "ohlc": ohlc,
595 "donchian": donchian,
596 "structural": structural,
597 "stop_loss": stop_loss,
598 "signals": signals
599 }
600
601
602@router.websocket("/ws")
603async def websocket_endpoint(websocket: WebSocket):
604 """WebSocket endpoint for real-time updates"""
605 await websocket.accept()
606 _ws_connections.append(websocket)
607
608 try:
609 # Send current status on connect
610 await websocket.send_json({
611 "type": "connected",
612 "config": get_config_manager().get_config().model_dump()
613 })
614
615 while True:
616 # Keep connection alive and handle incoming messages
617 try:
618 data = await asyncio.wait_for(websocket.receive_json(), timeout=30)
619
620 # Handle ping/pong
621 if data.get("type") == "ping":
622 await websocket.send_json({"type": "pong"})
623
624 except asyncio.TimeoutError:
625 # Send heartbeat
626 await websocket.send_json({"type": "heartbeat"})
627
628 except WebSocketDisconnect:
629 logger.info("WebSocket client disconnected")
630 except Exception as e:
631 logger.error(f"WebSocket error: {e}")
632 finally:
633 if websocket in _ws_connections:
634 _ws_connections.remove(websocket)
635
636
637async def broadcast_message(message: dict):
638 """Broadcast message to all connected WebSocket clients"""
639 disconnected = []
640 for ws in _ws_connections:
641 try:
642 await ws.send_json(message)
643 except Exception:
644 disconnected.append(ws)
645
646 # Clean up disconnected clients
647 for ws in disconnected:
648 if ws in _ws_connections:
649 _ws_connections.remove(ws)
650 