import pandas as pd
import numpy as np
from datetime import datetime, timedelta
from typing import Dict, List, Optional
import logging, time, sys, os, json
from pathlib import Path
from dataclasses import dataclass, field
from config import WaveStrategyConfig, DBConfig
from database import DatabaseManager
from wave_strategy import WaveRangeStrategy, WaveSignal
from virtual_exchange import VirtualExchange
from position_manager import PositionStatus
from report_generator import ReportGenerator  # 🔑 ИМПОРТ ГЕНЕРАТОРА ОТЧЕТОВ

logger = logging.getLogger(__name__)

@dataclass
class BacktestReport:
    """Отчет бэктеста с метриками, сделками и кривой доходности"""
    metrics: Dict
    trades: pd.DataFrame
    equity_curve: pd.DataFrame
    config_snapshot: Dict = field(default_factory=dict)
    timestamp: str = field(default_factory=lambda: datetime.now().isoformat())

    def to_dict(self) -> Dict:
        """Сериализация для JSON"""
        return {
            'metrics': self.metrics,
            'trades': self.trades.to_dict(orient='records') if not self.trades.empty else [],
            'equity_curve': self.equity_curve.to_dict(orient='records') if not self.equity_curve.empty else [],
            'config': self.config_snapshot,
            'timestamp': self.timestamp
        }


class Backtester:
    def __init__(self, config: WaveStrategyConfig, db: DatabaseManager):
        self.config = config
        self.db = db
        self.exchange = VirtualExchange(config)
        self.trades: List[Dict] = []
        self.equity_points: List[Dict] = []
        self._history_cache: Dict = {}

    def run(self, instruments: List[str], timeframes: List[str], start_date: str, end_date: str, save_report: bool = True) -> BacktestReport:
        logger.info(f"📈 Запуск бэктеста: {instruments} {timeframes} | {start_date} → {end_date}")
        
        all_bars = self._fetch_bars(instruments, timeframes, start_date, end_date)
        if all_bars.empty: raise ValueError("❌ Нет исторических данных")
            
        all_bars = all_bars.sort_values('timestamp').reset_index(drop=True)
        self._history_cache = {k: g.sort_values('timestamp').reset_index(drop=True) for (k, g) in all_bars.groupby(['instrument', 'timeframe'])}
            
        total = len(all_bars)
        initial_balance = self.exchange.balance
        start_time = time.time()
        last_update = start_time
        use_tty = sys.stdout.isatty()
        
        if use_tty: sys.stdout.write('\033[?25l')
        
        for idx in range(total):
            try:
                bar = all_bars.iloc[idx]
                if pd.isna(bar[['Close','High','Low']]).any(): continue
                    
                self._process_bar(bar, idx)
                
                trade_date = datetime.fromtimestamp(float(bar['timestamp'])).strftime('%Y-%m-%d')
                if not self.equity_points or self.equity_points[-1]['date'] != trade_date:
                    self.equity_points.append({
                        'date': trade_date, 'timestamp': float(bar['timestamp']), 
                        'balance': self.exchange.balance,
                        'open_positions': len(self.exchange.position_manager.get_open_positions()),
                        'cum_pnl': self.exchange.balance - initial_balance
                    })
            except Exception as e:
                logger.debug(f"⚠️ Ошибка на баре {idx}: {e}")
                continue
                
            now = time.time()
            if now - last_update >= 0.5 or idx == total - 1:
                elapsed = now - start_time
                pct = ((idx+1)/total)*100
                eta = (elapsed/(idx+1))*(total-idx-1) if idx+1 > 0 else 0
                bar_date = datetime.fromtimestamp(float(bar['timestamp'])).strftime('%Y-%m-%d %H:%M')
                
                # 🔑 ИСПРАВЛЕНО: \r в начале, вывод в stdout (отделено от логов stderr)
                status = f"\r🔄 [{pct:5.1f}%] {bar['instrument']}_{bar['timeframe']:<7} | {bar_date} | Bar {idx+1}/{total} | ⏱️ {time.strftime('%H:%M:%S', time.gmtime(elapsed))} | ETA: {time.strftime('%H:%M:%S', time.gmtime(int(eta)))} | 💰 ${self.exchange.balance:,.2f} | 📊 {len(self.exchange.position_manager.get_open_positions())} open "
                sys.stdout.write(status)
                sys.stdout.flush()
                last_update = now
                
        if use_tty: 
            sys.stdout.write('\033[?25h\n✅ Бэктест завершен.\n')
        else: 
            print('\n✅ Бэктест завершен.')
        sys.stdout.flush()
        
        self._close_all_positions(all_bars)
        
        metrics = self._calc_metrics(initial_balance)
        config_snapshot = {
            'atr_period': getattr(self.config, 'atr_period', None), 
            'sl_atr_multiplier': getattr(self.config, 'sl_atr_multiplier', None),
            'risk_reward_ratio': getattr(self.config, 'risk_reward_ratio', None), 
            'trail_enabled': getattr(self.config, 'trail_enabled', None),
            'max_risk_percent': getattr(self.config, 'max_risk_percent', None)
        }
        report = BacktestReport(metrics=metrics, trades=pd.DataFrame(self.trades), 
                                equity_curve=pd.DataFrame(self.equity_points), config_snapshot=config_snapshot)
        if save_report: self._export_report(report)
        self._print_summary(report)
        return report

    def generate_reports(self, report: BacktestReport):
        """🔑 Генерация расширенных отчетов и графиков"""
        if not getattr(self.config, 'generate_reports', True):
            return
            
        generator = ReportGenerator(output_dir='backtest_reports')
        generator.generate_full_report(
            report.metrics, 
            report.trades, 
            report.equity_curve, 
            report.config_snapshot
        )
    
    def _fetch_bars(self, instruments, timeframes, start_date, end_date):
        dfs, start_ts = [], int(datetime.strptime(start_date, '%Y-%m-%d').timestamp()) if start_date else None
        end_ts = int((datetime.strptime(end_date, '%Y-%m-%d') + timedelta(days=1)).timestamp()) if end_date else None
        
        for inst in instruments:
            for tf in timeframes:
                limit = max(5000, getattr(self.config, 'wave_lookback', 2000) * 2)
                df = self.db.fetch_ohlc_data(inst, tf, limit=limit)
                if df.empty: continue
                if start_ts: df = df[df['timestamp'] >= start_ts]
                if end_ts: df = df[df['timestamp'] < end_ts]
                if df.empty: continue
                df['instrument'], df['timeframe'] = inst, tf; dfs.append(df)
        return pd.concat(dfs, ignore_index=True) if dfs else pd.DataFrame()

    def _process_bar(self, bar, idx):
        inst, tf = bar['instrument'], bar['timeframe']
        high, low, close = float(bar['High']), float(bar['Low']), float(bar['Close'])
        
        lookback = getattr(self.config, 'wave_lookback', 2000)
        df_hist = self._get_history(inst, tf, bar['timestamp'], lookback=lookback)
        if df_hist.empty or len(df_hist) < self.config.atr_period + 10: return

        atr_val = self.exchange.strategy.calculate_atr(df_hist, self.config.atr_period).iloc[-1]
        if pd.isna(atr_val) or atr_val <= 0: return

        debug_mode = getattr(self.config, 'debug_signals', False)
        signal = self.exchange.strategy.analyze(df_hist, inst, tf, debug=debug_mode)
        
        if signal:
            pos = self.exchange.execute_signal(signal, current_bar_idx=idx)
            if pos:
                pos.entry_bar_idx = idx
                self.trades.append({'id': pos.id, 'instrument': inst, 'timeframe': tf, 'side': pos.side.value,
                                    'entry': pos.entry_price, 'sl': pos.stop_loss, 'tp': pos.take_profit,
                                    'atr': pos.atr_at_entry, 'progress': signal.wave_progress_pct,
                                    'exit': None, 'reason': None, 'pnl': None, 'opened_at': pos.opened_at.isoformat()})

        exits = self.exchange.update_price(inst, close, high, low, atr_val, idx)
        for ex in exits:
            for t in self.trades:
                if t['id'] == ex['id']: t.update({'exit': ex['exit_price'], 'reason': ex['reason'], 'pnl': ex['pnl'], 'exit_time': ex.get('timestamp')}); break

    def _get_history(self, inst, tf, current_ts, lookback):
        key = (inst, tf)
        if key not in self._history_cache: return pd.DataFrame()
        df = self._history_cache[key]
        if df.empty: return pd.DataFrame()
        idx = df['timestamp'].searchsorted(current_ts, side='right')
        if idx == 0: return pd.DataFrame()
        return df.iloc[max(0, idx - lookback):idx].copy()

    def _close_all_positions(self, df):
        if df.empty: return
        last_prices = df.groupby(['instrument', 'timeframe'])['Close'].last().to_dict()
        
        for pos in list(self.exchange.position_manager.get_open_positions()):
            key = (pos.instrument, pos.timeframe)
            if key not in last_prices: continue
            last_price = last_prices[key]
            pnl = self.exchange.position_manager.close_position(pos, last_price, 'backtest_end', 0.001)
            self.exchange.balance += pnl
            for t in self.trades:
                if t['id'] == pos.id: t.update({'exit': last_price, 'reason': 'backtest_end', 'pnl': pnl, 'exit_time': datetime.now().isoformat()}); break

    def _calc_metrics(self, initial_balance):
        trades_df = pd.DataFrame(self.trades)
        if trades_df.empty or 'pnl' not in trades_df.columns:
            return {'total_trades': 0, 'win_rate': 0.0, 'profit_factor': 0.0, 'total_pnl': 0.0, 'total_pnl_percent': 0.0, 'max_drawdown_pct': 0.0, 'sharpe_ratio': 0.0, 'expectancy': 0.0, 'avg_win': 0.0, 'avg_loss': 0.0}
        
        closed = trades_df[trades_df['pnl'].notna()].copy()
        if closed.empty: return {'total_trades': 0, 'win_rate': 0.0, 'profit_factor': 0.0, 'total_pnl': 0.0, 'total_pnl_percent': 0.0, 'max_drawdown_pct': 0.0, 'sharpe_ratio': 0.0, 'expectancy': 0.0}
        
        wins = closed[closed['pnl'] > 0]
        losses = closed[closed['pnl'] <= 0]
        
        total_trades = len(closed)
        win_rate = (len(wins) / total_trades) * 100 if total_trades > 0 else 0
        gross_profit = float(wins['pnl'].sum()) if not wins.empty else 0.0
        gross_loss = abs(float(losses['pnl'].sum())) if not losses.empty else 0.0
        profit_factor = gross_profit / gross_loss if gross_loss > 0 else (99.99 if gross_profit > 0 else 0)
        
        avg_win = float(wins['pnl'].mean()) if not wins.empty else 0.0
        avg_loss = float(losses['pnl'].mean()) if not losses.empty else 0.0
        expectancy = (win_rate/100 * avg_win) - ((1 - win_rate/100) * abs(avg_loss))
        
        eq_df = pd.DataFrame(self.equity_points)
        max_dd_pct = 0.0
        if not eq_df.empty and 'balance' in eq_df.columns:
            eq_df = eq_df.sort_values('timestamp' if 'timestamp' in eq_df.columns else 'date')
            eq_df['peak'] = eq_df['balance'].cummax()
            eq_df['drawdown'] = eq_df['balance'] - eq_df['peak']
            max_dd_pct = (eq_df['drawdown'].min() / initial_balance) * 100 if initial_balance > 0 else 0
        
        sharpe = 0.0
        if len(eq_df) > 10:
            eq_df['ret'] = eq_df['balance'].pct_change().dropna()
            if eq_df['ret'].std() > 1e-9: sharpe = (eq_df['ret'].mean() / eq_df['ret'].std()) * (365 ** 0.5)
        
        return {
            'total_trades': total_trades, 'win_rate': round(win_rate, 2), 'profit_factor': round(profit_factor, 2),
            'total_pnl': round(float(closed['pnl'].sum()), 2), 'total_pnl_percent': round((closed['pnl'].sum() / initial_balance) * 100, 2),
            'expectancy': round(expectancy, 2), 'avg_win': round(avg_win, 2), 'avg_loss': round(avg_loss, 2),
            'max_drawdown_pct': round(max_dd_pct, 2), 'sharpe_ratio': round(sharpe, 2),
            'best_trade': round(float(closed['pnl'].max()), 2), 'worst_trade': round(float(closed['pnl'].min()), 2)
        }

    def _export_report(self, report):
        out_dir = Path('backtest_results'); out_dir.mkdir(exist_ok=True)
        ts = datetime.now().strftime('%Y%m%d_%H%M%S')
        if not report.trades.empty: report.trades.to_csv(out_dir / f'trades_{ts}.csv', index=False, float_format='%.8f')
        if not report.equity_curve.empty: report.equity_curve.to_csv(out_dir / f'equity_{ts}.csv', index=False, float_format='%.2f')
        with open(out_dir / f'metrics_{ts}.json', 'w') as f: json.dump(report.to_dict(), f, indent=2, default=str)
        logger.info(f"📁 Отчёты сохранены в: {out_dir}")

    def _print_summary(self, report):
        m = report.metrics
        print("\n" + "="*80)
        print("📊 РЕЗУЛЬТАТЫ БЭКТЕСТА")
        print("="*80)
        print(f"Всего сделок: {m.get('total_trades', 0):>6} | Win Rate: {m.get('win_rate', 0):>5.2f}% | PF: {m.get('profit_factor', 0):>5.2f}")
        print(f"Total P&L: ${m.get('total_pnl', 0):>10,.2f} ({m.get('total_pnl_percent', 0):>6.2f}%) | DD: {m.get('max_drawdown_pct', 0):>5.2f}%")
        print(f"Expectancy: ${m.get('expectancy', 0):>8,.2f} | Sharpe: {m.get('sharpe_ratio', 0):>5.2f}")
        print("-"*80)
        if not report.trades.empty and 'pnl' in report.trades.columns:
            valid = report.trades.dropna(subset=['pnl'])
            if not valid.empty:
                print("🏆 Топ-3 прибыльных:"); print(valid.nlargest(3, 'pnl')[['instrument','side','pnl','reason']].to_string(index=False))
                print("📉 Топ-3 убыточных:"); print(valid.nsmallest(3, 'pnl')[['instrument','side','pnl','reason']].to_string(index=False))
        print("="*80 + "\n")