import pandas as pd
import numpy as np
from typing import Optional, List, Dict, Tuple
from dataclasses import dataclass, field
from datetime import datetime
from enum import Enum, auto
import logging

logger = logging.getLogger(__name__)

class SignalType(Enum):
    BUY = auto()
    SELL = auto()

@dataclass
class WaveSignal:
    instrument: str
    timeframe: str
    signal: SignalType
    price: float
    atr: float
    wave_type: str
    wave_progress_pct: float
    avg_wave_length_pct: float
    remaining_pct: float
    stop_loss: Optional[float] = None
    take_profit: Optional[float] = None
    confidence: float = 0.0
    reason: str = ""
    timestamp: datetime = field(default_factory=datetime.now)

    @property
    def risk_amount(self) -> float: return abs(self.price - self.stop_loss) if self.stop_loss else 0
    @property
    def reward_amount(self) -> float: return abs(self.take_profit - self.price) if self.take_profit else 0
    @property
    def rr_ratio(self) -> float: return self.reward_amount / self.risk_amount if self.risk_amount > 0 else 0

    def to_dict(self) -> Dict:
        return {k: v for k, v in vars(self).items() if not k.startswith('_') and k != 'timestamp'}

class WaveRangeStrategy:
    def __init__(self, config): self.config = config

    @staticmethod
    def calculate_atr(df: pd.DataFrame, period: int = 14) -> pd.Series:
        high, low, close = df['High'], df['Low'], df['Close']
        tr = pd.concat([high - low, (high - close.shift(1)).abs(), (low - close.shift(1)).abs()], axis=1).max(axis=1)
        return tr.rolling(window=period).mean()

    def detect_waves(self, df: pd.DataFrame) -> Tuple[List[Dict], Optional[Dict]]:
        """🔑 УПРОЩЁННЫЙ ДЕТЕКТОР: ищем только явные пики/впадины"""
        params = {
            'lookback': self.config.wave_lookback,
            'window': self.config.wave_window,
            'max_waves': self.config.max_waves_for_avg,
            'min_wave_distance_pct': self.config.min_wave_distance_pct,
        }
        df_work = df.iloc[-params['lookback']:].copy() if len(df) >= params['lookback'] else df.copy()
        current_price = df_work['Close'].iloc[-1]
        highs, lows = df_work['High'].values, df_work['Low'].values
        n = len(df_work)
        window = min(params['window'], max(3, n // 5))
        if n < window * 2 + 2: return [], None

        peaks, troughs = [], []
        for i in range(window, n - window):
            # 🔑 Строгая проверка: пик должен быть строго выше всех соседей
            if all(highs[i] > highs[i-j] and highs[i] > highs[i+j] for j in range(1, window+1)):
                peaks.append((i, highs[i]))
            if all(lows[i] < lows[i-j] and lows[i] < lows[i+j] for j in range(1, window+1)):
                troughs.append((i, lows[i]))

        extremes = sorted([(idx, 'peak', p) for idx, p in peaks] + [(idx, 'trough', p) for idx, p in troughs], key=lambda x: x[0])
        if len(extremes) < 2: return [], None

        waves = []
        for k in range(len(extremes) - 1):
            idx1, t1, p1 = extremes[k]
            idx2, t2, p2 = extremes[k+1]
            if (t1 == 'trough' and t2 == 'peak') or (t1 == 'peak' and t2 == 'trough'):
                length_pct = abs(p2 - p1) / p1 * 100 if p1 > 0 else 0
                if length_pct >= params['min_wave_distance_pct']:
                    waves.append({'type': 'up' if t1=='trough' else 'down', 'length_pct': length_pct, 'start': p1, 'end': p2})

        if not waves: return [], None

        # 🔑 Определяем текущую волну и прогресс
        last_type = extremes[-1][1]
        last_price = extremes[-1][2]
        current_wave_type = 'up' if last_type == 'trough' else 'down'
        distance_pct = abs(current_price - last_price) / current_price * 100 if current_price > 0 else 0

        # 🔑 Усредняем ТОЛЬКО по волнам того же типа
        same_type = [w for w in waves if w['type'] == current_wave_type]
        if same_type:
            last_waves = same_type[-params['max_waves']:]
            avg_len = sum(w['length_pct'] for w in last_waves) / len(last_waves)
        else:
            # Если нет исторических волн того же типа — берём консервативную оценку
            avg_len = distance_pct * 1.5

        remaining = max(0, avg_len - distance_pct)
        progress = (distance_pct / avg_len * 100) if avg_len > 0 else 100

        return waves, {
            'type': current_wave_type,
            'avg_wave_length_pct': avg_len,
            'remaining_pct': remaining,
            'progress_pct': progress,
            'last_extreme_price': last_price  # 🔑 Для структурного стопа
        }

    def calculate_levels(self, price: float, atr: float, signal_type: SignalType, last_extreme: float) -> Tuple[float, float]:
        """🔑 СТОП ЗА СТРУКТУРОЙ, а не просто ATR×mult"""
        if signal_type == SignalType.BUY:
            # Для BUY стоп должен быть НИЖЕ последнего минимума (extreme)
            sl = min(last_extreme * 0.998, price - atr * self.config.sl_atr_multiplier)  # 0.2% буфер
            tp = price + (price - sl) * self.config.risk_reward_ratio
        else:
            # Для SELL стоп должен быть ВЫШЕ последнего максимума
            sl = max(last_extreme * 1.002, price + atr * self.config.sl_atr_multiplier)
            tp = price - (sl - price) * self.config.risk_reward_ratio
        return round(sl, 2), round(tp, 2)  # 🔑 Округление до копеек для акций

    def analyze(self, df: pd.DataFrame, instrument: str, timeframe: str, debug: bool = False) -> Optional[WaveSignal]:
        min_len = max(self.config.atr_period + 30, 80)
        if len(df) < min_len: return None

        atr = self.calculate_atr(df, self.config.atr_period).iloc[-1]
        price = df['Close'].iloc[-1]
        if pd.isna(atr) or atr <= 0 or price <= 0: return None
        atr_pct = (atr / price) * 100

        # 🔑 ФИЛЬТР 1: Волатильность (не торгуем в «мёртвом» рынке)
        if atr_pct < self.config.min_volatility_pct:
            return None

        waves, wave_info = self.detect_waves(df)
        if not wave_info: return None

        # 🔑 ФИЛЬТР 2: Тренд (EMA50)
        if self.config.use_trend_filter:
            ema = df['Close'].ewm(span=self.config.trend_ema_period, adjust=False).mean().iloc[-1]
            if wave_info['type'] == 'up' and price < ema: return None
            if wave_info['type'] == 'down' and price > ema: return None

        signal_type = SignalType.BUY if wave_info['type'] == 'up' else SignalType.SELL
        
        # 🔑 ПРОВЕРКА РАЗРЕШЕНИЯ КОРОТКИХ ПОЗИЦИЙ
        if signal_type == SignalType.SELL and not getattr(self.config, 'allow_short_positions', True):
            if debug: logger.debug(f"🚫 {instrument}_{timeframe}: Короткие позиции отключены в конфиге")
            return None
            
        progress = wave_info['progress_pct']

        # 🔑 ФИЛЬТР 3: Прогресс волны (уже диапазон для акций)
        if progress < self.config.min_entry_progress_pct or progress > self.config.max_entry_progress_pct:
            return None

        # 🔑 ФИЛЬТР 4: Остаток хода (должен позволять достичь TP)
        required_move = atr_pct * self.config.sl_atr_multiplier * self.config.risk_reward_ratio
        if wave_info['remaining_pct'] < required_move:
            return None

        # 🔑 Расчёт уровней с привязкой к структуре
        sl, tp = self.calculate_levels(price, atr, signal_type, wave_info['last_extreme_price'])
        
        # 🔑 Дополнительная проверка: стоп не должен быть слишком близко (<0.3% от цены)
        if abs(price - sl) / price < 0.003:
            return None

        return WaveSignal(
            instrument=instrument, timeframe=timeframe, signal=signal_type,
            price=round(price, 2), atr=atr, wave_type=wave_info['type'],
            wave_progress_pct=progress, avg_wave_length_pct=wave_info['avg_wave_length_pct'],
            remaining_pct=wave_info['remaining_pct'], stop_loss=sl, take_profit=tp,
            confidence=0.75, reason=f"Волна {wave_info['type']} | Прогресс {progress:.1f}% | Структурный стоп"
        )

    def check_exit(self, position: Dict, close: float, high: float, low: float, wave_progress_pct: float = 0.0) -> Optional[Tuple[str, float]]:
        side, sl, tp = position['side'], position['stop_loss'], position['take_profit']
        # 🔑 Внутрибаровая проверка: сначала проверяем пробой стопа, затем тейка
        if side == 'BUY':
            if low <= sl: return 'stop_loss', sl
            if high >= tp: return 'take_profit', tp
        else:
            if high >= sl: return 'stop_loss', sl
            if low <= tp: return 'take_profit', tp
        if wave_progress_pct > 80: return 'wave_exhausted', close
        return None