import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
from matplotlib.gridspec import GridSpec
import seaborn as sns
from pathlib import Path
from datetime import datetime
from typing import Dict, Optional
import logging

logger = logging.getLogger(__name__)

class ReportGenerator:
    """Генератор профессиональных отчетов с графиками"""
    
    def __init__(self, output_dir: str = 'backtest_reports'):
        self.output_dir = Path(output_dir)
        self.output_dir.mkdir(parents=True, exist_ok=True)
        self.timestamp = datetime.now().strftime('%Y%m%d_%H%M%S')
        plt.style.use('seaborn-v0_8')
        sns.set_palette("husl")
        
    def generate_full_report(self, metrics: Dict, trades_df: pd.DataFrame, equity_df: pd.DataFrame, config: Dict):
        """Генерация полного отчета"""
        logger.info("📊 Генерация отчетов...")
        
        # 1. Текстовый вывод в консоль (улучшенный)
        self._print_console_report(metrics, trades_df, config)
        
        if trades_df.empty or equity_df.empty:
            logger.warning("⚠️ Нет данных для построения графиков")
            return
            
        # 2. Графики
        self._generate_charts(metrics, trades_df, equity_df)
        
        logger.info(f"✅ Отчеты сохранены в: {self.output_dir}")

    def _print_console_report(self, metrics: Dict, trades_df: pd.DataFrame, config: Dict):
        """Красивый вывод в консоль"""
        print("\n" + "="*100)
        print("📊 " + "РЕЗУЛЬТАТЫ БЭКТЕСТА".center(96) + " 📊")
        print("="*100)
        print(f"\n{'📈 ОБЩАЯ СТАТИСТИКА':<50} {'⚙️ ПАРАМЕТРЫ':<50}")
        print("-"*100)
        
        # Безопасный вывод метрик
        safe_val = lambda k, fmt: f"{metrics.get(k, 0):{fmt}}"
        print(f"{'Всего сделок:':<30} {safe_val('total_trades', '>19,')}")
        print(f"{'Win Rate:':<30} {safe_val('win_rate', '>18.2f')}%")
        print(f"{'Profit Factor:':<30} {safe_val('profit_factor', '>19.2f')}")
        print(f"{'Total P&L:':<30} ${safe_val('total_pnl', '>18,.2f')}")
        print(f"{'Expectancy:':<30} ${safe_val('expectancy', '>18,.2f')}")
        print(f"{'Max Drawdown:':<30} {safe_val('max_drawdown_pct', '>18.2f')}%")
        print(f"{'Sharpe Ratio:':<30} {safe_val('sharpe_ratio', '>19.2f')}")
        
        # Безопасный вывод конфига (защита от None)
        atr = config.get('atr_period')
        sl = config.get('sl_atr_multiplier')
        rr = config.get('risk_reward_ratio')
        trail = config.get('trail_enabled')
        
        print(f"\n{'ATR Period:':<30} {str(atr) if atr is not None else 'N/A':>19}")
        print(f"{'SL Multiplier:':<30} {f'{sl:.1f}x' if isinstance(sl, (int, float)) else 'N/A':>19}")
        print(f"{'Risk/Reward:':<30} {f'1:{rr}' if isinstance(rr, (int, float)) else 'N/A':<18}")
        print(f"{'Trail Enabled:':<30} {str(trail).lower() if trail is not None else 'N/A':<19}")

        if not trades_df.empty and 'pnl' in trades_df.columns:
            closed = trades_df[trades_df['pnl'].notna()]
            if not closed.empty:
                wins = closed[closed['pnl'] > 0]
                losses = closed[closed['pnl'] <= 0]
                print(f"\n{'📊 СТАТИСТИКА СДЕЛОК':<50}")
                print("-"*50)
                if not wins.empty: print(f"Прибыльных: {len(wins):>6} | Ср. выигрыш: ${wins['pnl'].mean():>10,.2f}")
                if not losses.empty: print(f"Убыточных: {len(losses):>6} | Ср. проигрыш: ${losses['pnl'].mean():>10,.2f}")
        print("\n" + "="*100)

    def _generate_charts(self, metrics, trades_df, equity_df):
        fig = plt.figure(figsize=(16, 12))
        gs = GridSpec(3, 2, figure=fig, hspace=0.3, wspace=0.25)
        
        self._plot_equity_curve(fig.add_subplot(gs[0, :]), equity_df)
        self._plot_drawdown(fig.add_subplot(gs[1, 0]), equity_df)
        self._plot_pnl_distribution(fig.add_subplot(gs[1, 1]), trades_df)
        self._plot_monthly_returns(fig.add_subplot(gs[2, 0]), equity_df)
        self._plot_win_loss_chart(fig.add_subplot(gs[2, 1]), trades_df)
        
        chart_path = self.output_dir / f'backtest_report_{self.timestamp}.png'
        plt.savefig(chart_path, dpi=150, bbox_inches='tight')
        plt.close()
        logger.info(f"📈 График сохранен: {chart_path}")

    def _plot_equity_curve(self, ax, equity_df):
        if 'balance' not in equity_df.columns: return
        eq = equity_df.sort_values('timestamp' if 'timestamp' in equity_df.columns else 'date')
        ax.plot(eq.index, eq['balance'], linewidth=2, color='#2E86AB')
        ax.fill_between(eq.index, eq['balance'], alpha=0.3, color='#2E86AB')
        if len(eq) > 0: ax.axhline(y=eq['balance'].iloc[0], color='gray', linestyle='--', alpha=0.5)
        ax.set_title('📈 Equity Curve', fontsize=14, fontweight='bold')
        ax.set_ylabel('Balance ($)'); ax.grid(True, alpha=0.3)
        
    def _plot_drawdown(self, ax, equity_df):
        if 'balance' not in equity_df.columns: return
        eq = equity_df.sort_values('timestamp' if 'timestamp' in equity_df.columns else 'date')
        eq['peak'] = eq['balance'].cummax()
        eq['drawdown'] = ((eq['balance'] - eq['peak']) / eq['peak'] * 100)
        ax.fill_between(eq.index, eq['drawdown'], 0, where=eq['drawdown'] < 0, interpolate=True, color='#E94F37', alpha=0.6)
        ax.set_title('📉 Drawdown', fontsize=14, fontweight='bold')
        ax.set_ylabel('Drawdown (%)'); ax.grid(True, alpha=0.3)
        
    def _plot_pnl_distribution(self, ax, trades_df):
        if trades_df.empty or 'pnl' not in trades_df.columns: return
        closed = trades_df[trades_df['pnl'].notna()]['pnl']
        if closed.empty: return
        ax.hist([closed[closed>0], closed[closed<=0]], bins=20, label=['Wins', 'Losses'], color=['#28A745', '#DC3545'], alpha=0.7, edgecolor='black')
        ax.set_title('💰 P&L Distribution', fontsize=14, fontweight='bold')
        ax.set_xlabel('P&L ($)'); ax.set_ylabel('Frequency'); ax.legend(); ax.grid(True, alpha=0.3)
        
    def _plot_monthly_returns(self, ax, equity_df):
        if 'balance' not in equity_df.columns or len(equity_df) < 2: return
        eq = equity_df.sort_values('timestamp' if 'timestamp' in equity_df.columns else 'date')
        eq['date'] = pd.to_datetime(eq['timestamp'] if 'timestamp' in eq.columns else eq['date'])
        eq['month'] = eq['date'].dt.to_period('M')
        monthly = eq.groupby('month')['balance'].last().pct_change() * 100
        colors = ['#28A745' if x > 0 else '#DC3545' for x in monthly.dropna()]
        ax.bar(range(len(monthly)), monthly.dropna(), color=colors, alpha=0.7)
        ax.set_title('📅 Monthly Returns', fontsize=14, fontweight='bold')
        ax.set_ylabel('Return (%)'); ax.grid(True, alpha=0.3, axis='y')
        
    def _plot_win_loss_chart(self, ax, trades_df):
        if trades_df.empty or 'pnl' not in trades_df.columns: return
        closed = trades_df[trades_df['pnl'].notna()]
        if closed.empty: return
        wins = len(closed[closed['pnl'] > 0])
        losses = len(closed[closed['pnl'] <= 0])
        total = wins + losses
        if total == 0: return
        sizes = [wins, losses]
        labels = [f'Wins\n{wins} ({wins/total*100:.1f}%)', f'Losses\n{losses} ({losses/total*100:.1f}%)']
        # 🔑 Исправлено: ax.pie возвращает кортеж из 3 элементов, не распаковываем
        ax.pie(sizes, labels=labels, colors=['#28A745', '#DC3545'], startangle=90, textprops={'fontsize': 12})
        ax.set_title('🎯 Win/Loss Ratio', fontsize=14, fontweight='bold')