# src/data/features.py
import numpy as np
import pandas as pd
from typing import Tuple
from config import BASE_HORIZON, MIN_HORIZON, MAX_HORIZON, SEQ_LEN, DEFAULT_SL_MODE, DEFAULT_SLIPPAGE_PCT, DEFAULT_HOLD_ON_TIMEOUT


# ============================================================
#  Индикаторы
# ============================================================

def calculate_atr(df: pd.DataFrame, period: int = 14) -> pd.Series:
    """Average True Range по Wilder."""
    high_low = df['High'] - df['Low']
    high_prev = (df['High'] - df['Close'].shift(1)).abs()
    low_prev = (df['Low'] - df['Close'].shift(1)).abs()
    true_range = pd.concat([high_low, high_prev, low_prev], axis=1).max(axis=1)
    return true_range.ewm(alpha=1 / period, adjust=False).mean()


def calculate_rsi(series: pd.Series, period: int = 14) -> pd.Series:
    """RSI, нормализованно."""
    delta = series.diff()
    gain = delta.where(delta > 0, 0.0)
    loss = -delta.where(delta < 0, 0.0)
    avg_gain = gain.rolling(period).mean()
    avg_loss = loss.rolling(period).mean()
    rs = avg_gain / (avg_loss + 1e-8)
    return 100.0 - (100.0 / (1.0 + rs))


def calculate_adx(df: pd.DataFrame, period: int = 14) -> pd.Series:
    """ADX в диапазоне [0, 1]."""
    high_diff = df['High'].diff()
    low_diff = -df['Low'].diff()
    plus_dm = np.where((high_diff > low_diff) & (high_diff > 0), high_diff, 0.0)
    minus_dm = np.where((low_diff > high_diff) & (low_diff > 0), low_diff, 0.0)
    tr = pd.concat([
        df['High'] - df['Low'],
        (df['High'] - df['Close'].shift(1)).abs(),
        (df['Low'] - df['Close'].shift(1)).abs()
    ], axis=1).max(axis=1)
    atr_smooth = tr.rolling(period).mean()
    plus_di = 100.0 * (pd.Series(plus_dm).rolling(period).mean() / (atr_smooth + 1e-8))
    minus_di = 100.0 * (pd.Series(minus_dm).rolling(period).mean() / (atr_smooth + 1e-8))
    dx = 100.0 * (plus_di - minus_di).abs() / ((plus_di + minus_di) + 1e-8)
    return dx.ewm(alpha=1 / period, adjust=False).mean() / 100.0


def calculate_bb_width(series: pd.Series, period: int = 20, std_dev: float = 2.0) -> pd.Series:
    """Bollinger Band Width."""
    sma = series.rolling(period).mean()
    std = series.rolling(period).std()
    upper = sma + std_dev * std
    lower = sma - std_dev * std
    return (upper - lower) / (sma + 1e-8)


# ============================================================
#  Векторизованная генерация меток (3 класса: TP / SL / HOLD)
# ============================================================

def generate_labels_dual_direction(
    df: pd.DataFrame,
    atr: pd.Series,
    rr_ratio: float,
    sl_mode: str | None = None,
    hold_on_timeout: bool | None = None,
    slippage_pct: float | None = None,
) -> Tuple[pd.Series, pd.Series]:
    """Генерирует метки LONG/SHORT: 1.0 = TP, 0.0 = SL, NaN = HOLD.

    Параметры
    ---------
    sl_mode : str
        "atr_from_close" — SL = entry ± ATR  (рекоменгуется, постоянный риск 1R)
        "bar_low_high"   — SL = low − ATR / high + ATR  (старое поведение)
    hold_on_timeout : bool
        True  → ни TP ни SL не сработали → NaN  (бар исключается из обучения)
        False → то же самое → 0.0  (консервативно: «нет профита = проигрыш»)
    slippage_pct : float
        Доля цены, добавляемая как проскальзывание при входе.
        LONG:  entry = close × (1 + slippage)  → покупаем дороже
        SHORT: entry = close × (1 − slippage)  → продаём дешевле
    """
    _sl_mode = sl_mode or DEFAULT_SL_MODE
    _hold = hold_on_timeout if hold_on_timeout is not None else DEFAULT_HOLD_ON_TIMEOUT
    _slippage = slippage_pct if slippage_pct is not None else DEFAULT_SLIPPAGE_PCT
    n = len(df)
    nan_labels = pd.Series(np.full(n, np.nan), index=df.index)
    if n < MIN_HORIZON + SEQ_LEN:
        return nan_labels.copy(), nan_labels.copy()

    closes = df['Close'].values.astype(np.float64)
    highs = df['High'].values.astype(np.float64)
    lows = df['Low'].values.astype(np.float64)
    atr_vals = atr.values.astype(np.float64)

    # --- ref_vol: медиана волатильности на «чистом» участке ----------
    vol_ratio = atr_vals / (closes + 1e-8)
    valid_start = 20
    valid_end = max(MIN_HORIZON + 20, n - MAX_HORIZON - 1)
    if valid_end <= valid_start:
        return nan_labels.copy(), nan_labels.copy()

    valid_slice = vol_ratio[valid_start:valid_end]
    valid_vals = valid_slice[~np.isnan(valid_slice)]
    ref_vol = float(np.nanpercentile(valid_vals, 50)) if len(valid_vals) > 0 else 1e-8

    # --- dynamic horizons: высокая волатильность → короче горизонт ---
    dynamic_horizons = np.clip(
        np.round(BASE_HORIZON * (ref_vol / (vol_ratio + 1e-8))),
        MIN_HORIZON, MAX_HORIZON
    ).astype(int)

    # ================================================================
    #  ВЕКТОРИЗОВАННЫЙ РАСЧЁТ
    # ================================================================
    H = MAX_HORIZON                       # максимальное окно заглядывания
    valid_indices = np.arange(valid_start, valid_end)
    n_valid = len(valid_indices)
    if n_valid == 0:
        return nan_labels.copy(), nan_labels.copy()

    # --- Entry с проскальзыванием -------------------------------------
    close_entry = closes[valid_indices]
    atr_at_entry = atr_vals[valid_indices]

    entry_long = close_entry * (1.0 + _slippage)   # worse for LONG
    entry_short = close_entry * (1.0 - _slippage)   # worse for SHORT

    # --- SL / TP levels ------
    if _sl_mode == "bar_low_high":
        sl_long = lows[valid_indices] - atr_at_entry
        sl_short = highs[valid_indices] + atr_at_entry
    else:  # atr_from_close
        sl_long = close_entry - atr_at_entry
        sl_short = close_entry + atr_at_entry

    # Риск измеряется от entry (с учётом slippage)
    risk_long = entry_long - sl_long
    risk_short = sl_short - entry_short

    tp_long = entry_long + rr_ratio * risk_long
    tp_short = entry_short - rr_ratio * risk_short

    # --- 2D-матрицы будущих цен: (n_valid, H) ------------------------
    # row_offsets[i, s] = valid_indices[i] + 1 + s
    row_offsets = valid_indices[:, None] + np.arange(1, H + 1)[None, :]

    future_highs_2d = np.take(highs, row_offsets, mode='clip')
    future_lows_2d = np.take(lows, row_offsets, mode='clip')

    # --- Маска валидных шагов: s < horizon[i] ------------------------
    horizons = dynamic_horizons[valid_indices]          # shape (n_valid,)
    step_indices = np.arange(1, H + 1)[None, :]          # shape (1, H)
    valid_mask = step_indices <= horizons[:, None]       # shape (n_valid, H)

    # Не-валидные шаги: подставляем значения, которые НЕ сработают
    future_lows_masked = np.where(valid_mask, future_lows_2d, np.inf)
    future_highs_masked = np.where(valid_mask, future_highs_2d, -np.inf)

    # --- LONG: первый hit SL и TP -------------------------------------
    sl_hit_long_mask = future_lows_masked <= sl_long[:, None]
    tp_hit_long_mask = future_highs_masked >= tp_long[:, None]

    any_sl_long = np.any(sl_hit_long_mask, axis=1)
    any_tp_long = np.any(tp_hit_long_mask, axis=1)

    hit_sl_long = np.where(any_sl_long, np.argmax(sl_hit_long_mask, axis=1), -1)
    hit_tp_long = np.where(any_tp_long, np.argmax(tp_hit_long_mask, axis=1), -1)

    # --- SHORT: первый hit SL и TP ------------------------------------
    sl_hit_short_mask = future_highs_masked >= sl_short[:, None]
    tp_hit_short_mask = future_lows_masked <= tp_short[:, None]

    any_sl_short = np.any(sl_hit_short_mask, axis=1)
    any_tp_short = np.any(tp_hit_short_mask, axis=1)

    hit_sl_short = np.where(any_sl_short, np.argmax(sl_hit_short_mask, axis=1), -1)
    hit_tp_short = np.where(any_tp_short, np.argmax(tp_hit_short_mask, axis=1), -1)

    # --- Разрешение меток ----------------------------------------------
    labels_long = _resolve_labels(hit_tp_long, hit_sl_long, _hold)
    labels_short = _resolve_labels(hit_tp_short, hit_sl_short, _hold)

    # --- Сборка выходных серий -----------------------------------------
    out_long = np.full(n, np.nan)
    out_short = np.full(n, np.nan)
    out_long[valid_indices] = labels_long
    out_short[valid_indices] = labels_short

    return pd.Series(out_long, index=df.index), pd.Series(out_short, index=df.index)


# ---------------------------------------------------------------
#  Вспомогательные функции
# ---------------------------------------------------------------

def _resolve_labels(hit_tp: np.ndarray, hit_sl: np.ndarray,
                    hold_on_timeout: bool) -> np.ndarray:
    """Векторизованное разрешение меток.

    1.0 — TP сработал раньше (или одновременно) с SL
    0.0 — SL сработал раньше
    NaN — ни TP ни SL (HOLD / таймаут)
    """
    n = len(hit_tp)
    labels = np.full(n, np.nan)

    both_hit = (hit_tp != -1) & (hit_sl != -1)
    only_tp = (hit_tp != -1) & (hit_sl == -1)
    only_sl = (hit_tp == -1) & (hit_sl != -1)
    neither = (hit_tp == -1) & (hit_sl == -1)

    labels[both_hit & (hit_tp <= hit_sl)] = 1.0
    labels[both_hit & (hit_tp > hit_sl)] = 0.0
    labels[only_tp] = 1.0
    labels[only_sl] = 0.0
    labels[neither] = np.nan if hold_on_timeout else 0.0

    return labels


# ============================================================
#  Feature engineering
# ============================================================

def engineer_features(df: pd.DataFrame, window: int = 20,
                      rr_ratio: float = 3.0) -> pd.DataFrame:
    """Расчёт инвариантных признаков и двухнаправленных меток."""
    w = window
    df = df.copy()

    # --- Доходности ---
    df['Ret_1'] = np.log(df['Close'] / (df['Close'].shift(1) + 1e-8))
    df['Ret_5'] = np.log(df['Close'] / (df['Close'].shift(5) + 1e-8))
    df['Ret_20'] = np.log(df['Close'] / (df['Close'].shift(20) + 1e-8))

    # --- ATR и волатильность ---
    df['ATR'] = calculate_atr(df, 14)
    df['Vol_Ratio'] = df['ATR'] / (df['Close'] + 1e-8)
    df['Vol_Spread'] = np.log(df['Volume'] / (df['Volume'].rolling(w).mean() + 1e-8))

    # --- Осцилляторы ---
    df['RSI_14'] = calculate_rsi(df['Close'], 14) / 100.0
    low_w = df['Low'].rolling(w).min()
    high_w = df['High'].rolling(w).max()
    df['Stoch_K'] = (df['Close'] - low_w) / (high_w - low_w + 1e-8)
    df['Price_Pos'] = df['Stoch_K'].copy()

    ema9 = df['Close'].ewm(span=9, adjust=False).mean()
    ema21 = df['Close'].ewm(span=21, adjust=False).mean()
    df['EMA_Ratio'] = ema9 / (ema21 + 1e-8)

    # --- Статистика доходностей ---
    df['Ret_Std'] = df['Ret_1'].rolling(w).std()
    df['Ret_Skew'] = df['Ret_1'].rolling(w).skew()
    df['Ret_Kurt'] = df['Ret_1'].rolling(w).kurt()

    # --- Тренд ---
    df['ADX_14'] = calculate_adx(df, 14)
    df['BB_Width'] = calculate_bb_width(df['Close'], 20)

    # --- Labels LONG / SHORT (3 classes: TP=1, SL=0, HOLD=NaN) ---
    df['Label_Long'], df['Label_Short'] = generate_labels_dual_direction(
        df, df['ATR'], rr_ratio,
    )

    # --- Очистка NaN ---
    price_cols = {'Open', 'High', 'Low', 'Close', 'Volume'}
    num_cols = [c for c in df.select_dtypes(include=[np.number]).columns
                if c not in price_cols]
    df[num_cols] = df[num_cols].fillna(0.0).clip(-10, 10)

    # Удаляем строки без меток (HOLD-бары + NaN)
    df = df.dropna(subset=['Label_Long', 'Label_Short']).reset_index(drop=True)
    return df
