import itertools
import multiprocessing as mp
import time
import logging
from typing import Dict, List, Optional
from dataclasses import dataclass, field, asdict
from datetime import datetime
from config import WaveStrategyConfig, DBConfig, OptimizationConfig
from database import DatabaseManager
from backtester import Backtester
from param_store import ParameterStore, OptimizedParams

logger = logging.getLogger(__name__)

@dataclass
class OptResult:
    params: Dict
    metrics: Dict
    fitness: float
    exec_time: float
    status: str = 'success'

class ParameterOptimizer:
    def __init__(self, config: WaveStrategyConfig, db: DatabaseManager, opt_config: OptimizationConfig = None):
        self.config = config
        self.db = db
        self.opt_config = opt_config or OptimizationConfig()
        self.store = ParameterStore()

    def optimize(self, instrument: str, timeframe: str, start_date: str, end_date: str) -> Optional[OptimizedParams]:
        logger.info(f"🔍 Оптимизация: {instrument}_{timeframe}")
        grid = self._gen_grid()
        logger.info(f"📊 Сетка: {len(grid)} комбинаций")
        
        # Передача DBConfig в дочерние процессы
        db_cfg_dict = asdict(self.db.config)
        
        tasks = [{
            'instrument': instrument,
            'timeframe': timeframe,
            'params': p,
            'start': start_date,
            'end': end_date,
            'opt_cfg': asdict(self.opt_config),
            'base_cfg': asdict(self.config),
            'db_cfg': db_cfg_dict
        } for p in grid]
        
        start = time.time()
        if self.opt_config.n_jobs == 1:
            results = [self._run(t) for t in tasks]
        else:
            n_workers = mp.cpu_count() if self.opt_config.n_jobs == -1 else self.opt_config.n_jobs
            logger.info(f"🔄 Параллельно: {n_workers} процессов")
            with mp.Pool(processes=n_workers) as pool:
                results = pool.map(self._run, tasks)
                
        logger.info(f"✅ За {time.time()-start:.1f}с")
        valid = [r for r in results if r.status == 'success']
        if not valid:
            logger.error("❌ Нет валидных результатов")
            return None
        
        best = max(valid, key=lambda r: r.fitness)
        opt = OptimizedParams(
            instrument=instrument, timeframe=timeframe,
            atr_period=best.params['atr_period'],
            sl_atr_multiplier=best.params['sl_atr_multiplier'],
            risk_reward_ratio=best.params['risk_reward_ratio'],
            trail_distance_atr=best.params['trail_distance_atr'],
            fitness_score=best.fitness,
            total_trades=best.metrics.get('total_trades', 0),
            win_rate=best.metrics.get('win_rate', 0),
            profit_factor=best.metrics.get('profit_factor', 0),
            max_drawdown_pct=best.metrics.get('max_drawdown_pct', 0),
            backtest_start=start_date,
            backtest_end=end_date
        )
        self.store.save(opt)
        self._print_summary(opt, valid)
        return opt

    def _gen_grid(self) -> List[Dict]:
        return [{
            'atr_period': a,
            'sl_atr_multiplier': s,
            'risk_reward_ratio': r,
            'trail_distance_atr': t
        } for a in self.opt_config.atr_periods
          for s in self.opt_config.sl_multipliers
          for r in self.opt_config.rr_ratios
          for t in self.opt_config.trail_distances]

    @staticmethod
    def _run(task: Dict) -> OptResult:
        import logging
        logging.getLogger().setLevel(logging.WARNING)
        
        inst = task['instrument']
        tf = task['timeframe']
        params = task['params']
        start = task['start']
        end = task['end']
        opt_cfg = OptimizationConfig(**task['opt_cfg'])
        base_cfg = WaveStrategyConfig(**task['base_cfg'])
        db_cfg = DBConfig(**task['db_cfg'])
        
        try:
            cfg = WaveStrategyConfig()
            for k, v in asdict(base_cfg).items():
                if hasattr(cfg, k): 
                    setattr(cfg, k, v)
                    
            cfg.atr_period = params['atr_period']
            cfg.sl_atr_multiplier = params['sl_atr_multiplier']
            cfg.risk_reward_ratio = params['risk_reward_ratio']
            cfg.trail_distance_atr = params['trail_distance_atr']
            
            db = DatabaseManager(db_cfg, pool_size=1)
            bt = Backtester(cfg, db)
            
            # Walk-Forward логика (упрощенная)
            if opt_cfg.walk_forward:
                report = bt.run([inst], [tf], start, end)
            else:
                report = bt.run([inst], [tf], start, end)
                
            fitness = ParameterOptimizer._fitness(report.metrics, opt_cfg.fitness_function)
            
            if report.metrics.get('total_trades', 0) < opt_cfg.min_trades:
                return OptResult(params, report.metrics, -1, 0, 'filtered_trades')
            if report.metrics.get('win_rate', 0) < opt_cfg.min_win_rate:
                return OptResult(params, report.metrics, -1, 0, 'filtered_wr')
            if report.metrics.get('max_drawdown_pct', 100) > opt_cfg.max_drawdown_pct:
                return OptResult(params, report.metrics, -1, 0, 'filtered_dd')
            
            return OptResult(params, report.metrics, fitness, 0, 'success') 
            
        except Exception as e:
            return OptResult(params, {}, -1, 0, f'error:{str(e)}')

    @staticmethod
    def _fitness(metrics: Dict, method: str) -> float:
        if not metrics or 'error' in str(metrics):
            return -1.0
        if method == 'sharpe':
            return metrics.get('sharpe_ratio', 0)
        if method == 'profit_factor':
            pf = metrics.get('profit_factor', 0)
            return pf if pf != float('inf') else 10.0
            
        pf = min(metrics.get('profit_factor', 0), 10)
        sharpe = metrics.get('sharpe_ratio', 0)
        wr = metrics.get('win_rate', 0) / 100
        dd = metrics.get('max_drawdown_pct', 100) / 100
        return round(0.4*pf + 0.3*sharpe + 0.2*wr - 0.1*dd, 3)

    def _print_summary(self, best: OptimizedParams, results: List[OptResult]):
        print(f"\n🏆 {best.instrument}_{best.timeframe}")
        print(f"ATR={best.atr_period} | SL={best.sl_atr_multiplier}× | R:R=1:{best.risk_reward_ratio}")
        print(f"⭐ Fitness={best.fitness_score:.3f} | Win={best.win_rate:.1f}% | PF={best.profit_factor:.2f}")

    def optimize_all(self, instruments: List[str], timeframes: List[str], start: str, end: str):
        results = {}
        total = len(instruments) * len(timeframes)
        for i, inst in enumerate(instruments, 1):
            for tf in timeframes:
                logger.info(f"\n🔄 [{i}/{total}] {inst}_{tf}")
                results[f"{inst}_{tf}"] = self.optimize(inst, tf, start, end)
        return results