Team Ai
Apppublic

javitechjkd/backtestingv2

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
routes.py650 linesDownload Raw Back to api
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