import numpy as np
import pandas as pd


def _scan_outcome(
    closes: np.ndarray,
    highs: np.ndarray,
    lows: np.ndarray,
    atr_values: np.ndarray,
    atr_mult_sl: float,
    atr_mult_tp: float,
    max_bars: int,
    direction: int,
) -> np.ndarray:
    n = len(closes)
    outcomes = np.zeros(n, dtype=np.int32)

    for i in range(n - 1):
        if np.isnan(atr_values[i]) or atr_values[i] == 0:
            continue
        entry = closes[i]
        atr = atr_values[i]

        if direction == 1:
            tp = entry + atr * atr_mult_tp
            sl = entry - atr * atr_mult_sl
        else:
            tp = entry - atr * atr_mult_tp
            sl = entry + atr * atr_mult_sl

        limit = min(n, i + max_bars + 1)
        for j in range(i + 1, limit):
            if direction == 1:
                tp_hit = highs[j] >= tp
                sl_hit = lows[j] <= sl
            else:
                tp_hit = lows[j] <= tp
                sl_hit = highs[j] >= sl

            if tp_hit:
                outcomes[i] = 1
                break
            if sl_hit:
                break

    return outcomes


def compute_trade_outcome(
    df: pd.DataFrame,
    atr_mult_sl: float = 1.5,
    atr_mult_tp: float = 3.0,
    max_bars: int = 100,
) -> pd.DataFrame:
    df = df.copy()
    if 'atr' not in df.columns:
        raise ValueError("add_atr() must be called before compute_trade_outcome()")

    closes = df['Close'].values
    highs = df['High'].values
    lows = df['Low'].values
    atr_values = df['atr'].values
    n = len(df)

    outcomes = np.zeros(n, dtype=np.int32)

    for i in range(n - 1):
        if np.isnan(atr_values[i]) or atr_values[i] == 0:
            continue
        entry = closes[i]
        tp = entry + atr_values[i] * atr_mult_tp
        sl = entry - atr_values[i] * atr_mult_sl

        limit = min(n, i + max_bars + 1)
        for j in range(i + 1, limit):
            if highs[j] >= tp:
                outcomes[i] = 1
                break
            if lows[j] <= sl:
                outcomes[i] = 2
                break

    df['outcome'] = outcomes
    return df


def compute_dual_outcomes(
    df: pd.DataFrame,
    atr_mult_sl: float = 1.5,
    atr_mult_tp: float = 3.0,
    max_bars: int = 100,
) -> pd.DataFrame:
    df = df.copy()
    if 'atr' not in df.columns:
        raise ValueError("add_atr() must be called before compute_dual_outcomes()")

    arr = (
        df['Close'].values, df['High'].values, df['Low'].values,
        df['atr'].values,
    )
    df['outcome_long'] = _scan_outcome(*arr, atr_mult_sl, atr_mult_tp, max_bars, direction=1)
    df['outcome_short'] = _scan_outcome(*arr, atr_mult_sl, atr_mult_tp, max_bars, direction=-1)

    # Rolling WR использует упрощённый outcome с меньшим max_bars (30),
    # чтобы разрешённые строки были ближе к настоящему моменту
    df['outcome_long_fast'] = _scan_outcome(*arr, atr_mult_sl, atr_mult_tp, 30, direction=1)
    df['outcome_short_fast'] = _scan_outcome(*arr, atr_mult_sl, atr_mult_tp, 30, direction=-1)
    return df


def _find_valid_peak_high(
    high: np.ndarray,
    low: np.ndarray,
    start: int,
    end: int,
) -> float:
    """
    Ищет ближайший валидный пик в окне [start, end).
    
    Критерии пика:
      1. High[j] — локальный максимум (High[j] >= High[j-1] и High[j] >= High[j+1])
      2. Пик центрирован: в симметричном окне [j-r : j+r], где r = min(j-start, end-1-j),
         нет High выше, чем High[j] (т.е. j — истинная вершина на этом участке графика)
      3. Выбирается первый (ближайший к start) валидный пик
    
    Если валидный пик не найден — возвращает глобальный максимум окна (фолбэк).
    """
    # Сначала проверяем глобальный максимум
    global_max = np.max(high[start:end])
    argmax = start + np.argmax(high[start:end])
    
    # Проверяем центрированность глобального максимума
    left_dist = argmax - start
    right_dist = end - 1 - argmax
    radius = min(left_dist, right_dist)
    if radius >= 2:
        check_l = argmax - radius
        check_r = argmax + radius + 1
        if high[argmax] >= np.max(high[check_l:check_r]):
            return global_max  # глобальный максимум — валидный центрированный пик
    
    # Ищем первый валидный центрированный пик, сканируя слева направо
    for j in range(start + 1, end - 1):
        # Проверка: локальный максимум
        if high[j] >= high[j - 1] and high[j] >= high[j + 1]:
            left_d = j - start
            right_d = end - 1 - j
            r = min(left_d, right_d)
            if r >= 2:
                check_l = j - r
                check_r = j + r + 1
                if high[j] >= np.max(high[check_l:check_r]):
                    return high[j]  # первый валидный пик
    
    # Фолбэк: глобальный максимум
    return global_max


def _find_valid_trough_low(
    high: np.ndarray,
    low: np.ndarray,
    start: int,
    end: int,
) -> float:
    """
    Ищет ближайшую валидную впадину в окне [start, end).
    
    Критерии впадины:
      1. Low[j] — локальный минимум (Low[j] <= Low[j-1] и Low[j] <= Low[j+1])
      2. Впадина центрирована: в симметричном окне [j-r : j+r], где r = min(j-start, end-1-j),
         нет Low ниже, чем Low[j]
      3. Выбирается первая (ближайшая к start) валидная впадина
    
    Если валидная впадина не найдена — возвращает глобальный минимум окна (фолбэк).
    """
    # Сначала проверяем глобальный минимум
    global_min = np.min(low[start:end])
    argmin = start + np.argmin(low[start:end])
    
    left_dist = argmin - start
    right_dist = end - 1 - argmin
    radius = min(left_dist, right_dist)
    if radius >= 2:
        check_l = argmin - radius
        check_r = argmin + radius + 1
        if low[argmin] <= np.min(low[check_l:check_r]):
            return global_min  # глобальный минимум — валидная центрированная впадина
    
    # Ищем первую валидную центрированную впадину, сканируя слева направо
    for j in range(start + 1, end - 1):
        if low[j] <= low[j - 1] and low[j] <= low[j + 1]:
            left_d = j - start
            right_d = end - 1 - j
            r = min(left_d, right_d)
            if r >= 2:
                check_l = j - r
                check_r = j + r + 1
                if low[j] <= np.min(low[check_l:check_r]):
                    return low[j]  # первая валидная впадина
    
    # Фолбэк: глобальный минимум
    return global_min


def compute_peak_trough_targets(
    df: pd.DataFrame,
    horizon: int = 100,
) -> pd.DataFrame:
    """
    Для каждой свечи находит ближайший ВАЛИДНЫЙ пик и впадину
    в следующих `horizon` барах.
    
    Пик/впадина считаются валидными, если:
      - Это локальный экстремум (выше/ниже соседей)
      - Центрирован: в симметричном окне [pos-r : pos+r] нет точек выше/ниже
    
    Возвращает регрессионные цели:
      pct_to_peak:    High[пик] / Close[i] - 1  — потенциал роста (+, %)
      pct_to_trough:  Low[впадина] / Close[i] - 1  — потенциал падения (-, %)
    """
    df = df.copy()
    close = df['Close'].values
    high = df['High'].values
    low = df['Low'].values
    n = len(df)

    pct_to_peak = np.full(n, np.nan, dtype=np.float64)
    pct_to_trough = np.full(n, np.nan, dtype=np.float64)

    for i in range(n - horizon):
        start = i + 1
        end = i + 1 + horizon

        peak_high = _find_valid_peak_high(high, low, start, end)
        trough_low = _find_valid_trough_low(high, low, start, end)

        pct_to_peak[i] = peak_high / close[i] - 1.0
        pct_to_trough[i] = trough_low / close[i] - 1.0

    df['pct_to_peak'] = pct_to_peak
    df['pct_to_trough'] = pct_to_trough
    return df


def outcome_distribution(df: pd.DataFrame) -> pd.Series:
    return df['outcome'].value_counts().sort_index()
