import talib
import numpy as np
import pandas as pd


def add_ema(df: pd.DataFrame, period: int, col: str = 'Close') -> pd.Series:
    return talib.EMA(df[col].values, timeperiod=period)


def add_rsi(df: pd.DataFrame, period: int = 14) -> pd.Series:
    return talib.RSI(df['Close'].values, timeperiod=period)


def add_macd(df: pd.DataFrame, fast: int = 12, slow: int = 26, signal: int = 9):
    macd, signal_line, hist = talib.MACD(df['Close'].values, fastperiod=fast, slowperiod=slow, signalperiod=signal)
    return macd, signal_line, hist


def add_bbands(df: pd.DataFrame, period: int = 20, nbdev: float = 2.0):
    upper, middle, lower = talib.BBANDS(df['Close'].values, timeperiod=period, nbdevup=nbdev, nbdevdn=nbdev, matype=0)
    return upper, middle, lower


def add_atr(df: pd.DataFrame, period: int = 14) -> pd.Series:
    return talib.ATR(df['High'].values, df['Low'].values, df['Close'].values, timeperiod=period)


def add_stochastic(df: pd.DataFrame, k_period: int = 5, d_period: int = 3):
    slowk, slowd = talib.STOCH(df['High'].values, df['Low'].values, df['Close'].values,
                               fastk_period=k_period, slowk_period=d_period, slowk_matype=0,
                               slowd_period=d_period, slowd_matype=0)
    return slowk, slowd


def add_candle_patterns(df: pd.DataFrame) -> dict:
    return {
        'doji': talib.CDLDOJI(df['Open'].values, df['High'].values, df['Low'].values, df['Close'].values),
        'hammer': talib.CDLHAMMER(df['Open'].values, df['High'].values, df['Low'].values, df['Close'].values),
        'shooting_star': talib.CDLSHOOTINGSTAR(df['Open'].values, df['High'].values, df['Low'].values, df['Close'].values),
        'engulfing': talib.CDLENGULFING(df['Open'].values, df['High'].values, df['Low'].values, df['Close'].values),
    }


def detect_rsi_divergence(df: pd.DataFrame, rsi_col: str = 'RSI14', lookback: int = 20,
                          out_bull: str = 'rsi_bull_div', out_bear: str = 'rsi_bear_div') -> pd.DataFrame:
    """
    Detect RSI divergence over lookback window.
    Writes to out_bull (0/1) and out_bear (0/1) columns.
    """
    rsi = df[rsi_col].values
    close = df['Close'].values
    n = len(df)

    bull_div = np.zeros(n, dtype=int)
    bear_div = np.zeros(n, dtype=int)

    for i in range(lookback + 5, n):
        window_rsi = rsi[i - lookback:i + 1]
        window_close = close[i - lookback:i + 1]

        # Bullish divergence: price lower low, RSI higher low
        price_min_idx = np.argmin(window_close)
        if price_min_idx == len(window_close) - 1:
            prev_price_min = window_close[:price_min_idx].min()
            prev_price_min_idx = np.argmin(window_close[:price_min_idx])
            prev_rsi_min = window_rsi[prev_price_min_idx]
            cur_rsi_val = rsi[i]
            if window_close[-1] < prev_price_min and cur_rsi_val > prev_rsi_min and cur_rsi_val < 50:
                bull_div[i] = 1

        # Bearish divergence: price higher high, RSI lower high
        price_max_idx = np.argmax(window_close)
        if price_max_idx == len(window_close) - 1:
            prev_price_max = window_close[:price_max_idx].max()
            prev_price_max_idx = np.argmax(window_close[:price_max_idx])
            prev_rsi_max = window_rsi[prev_price_max_idx]
            cur_rsi_val = rsi[i]
            if window_close[-1] > prev_price_max and cur_rsi_val < prev_rsi_max and cur_rsi_val > 50:
                bear_div[i] = 1

    df[out_bull] = bull_div
    df[out_bear] = bear_div
    return df


def detect_sr_levels(df: pd.DataFrame, lookback: int = 20,
                     out_res: str = 'sr_resistance', out_sup: str = 'sr_support') -> pd.DataFrame:
    """
    Compute rolling support/resistance levels from swing highs/lows.
    Writes to out_res and out_sup columns.
    """
    n = len(df)
    resistance = np.full(n, np.nan)
    support = np.full(n, np.nan)

    for i in range(lookback, n):
        high_window = df['High'].iloc[i - lookback:i]
        low_window = df['Low'].iloc[i - lookback:i]
        resistance[i] = high_window.max()
        support[i] = low_window.min()

    df[out_res] = resistance
    df[out_sup] = support
    return df


def compute_all_indicators(df: pd.DataFrame) -> pd.DataFrame:
    # --- EMA variants ---
    df['EMA5'] = add_ema(df, 5)
    df['EMA13'] = add_ema(df, 13)
    df['EMA9'] = add_ema(df, 9)
    df['EMA21'] = add_ema(df, 21)

    # --- RSI ---
    df['RSI14'] = add_rsi(df, 14)

    # --- ATR + median ---
    df['ATR14'] = add_atr(df, 14)
    df['ATR14_median_50'] = df['ATR14'].rolling(50).median()

    # --- MACD (standard 12/26/9) ---
    macd, signal_line, hist = add_macd(df)
    df['MACD'] = macd
    df['MACD_signal'] = signal_line
    df['MACD_hist'] = hist

    # --- BB variants ---
    # BB(20, 2) — original
    bb_u20, bb_m20, bb_l20 = add_bbands(df, period=20, nbdev=2)
    df['BB_upper_20_2'] = bb_u20
    df['BB_middle_20_2'] = bb_m20
    df['BB_lower_20_2'] = bb_l20
    # BB(14, 2.5) — wider bands, shorter window → more touches
    bb_u14_25, bb_m14_25, bb_l14_25 = add_bbands(df, period=14, nbdev=2)
    df['BB_upper_14_25'] = bb_u14_25
    df['BB_middle_14_25'] = bb_m14_25
    df['BB_lower_14_25'] = bb_l14_25
    # BB(14, 2) — shorter window, original std
    bb_u14, bb_m14, bb_l14 = add_bbands(df, period=14, nbdev=2)
    df['BB_upper_14_2'] = bb_u14
    df['BB_middle_14_2'] = bb_m14
    df['BB_lower_14_2'] = bb_l14
    # BB(10, 2) — even shorter window
    bb_u10, bb_m10, bb_l10 = add_bbands(df, period=10, nbdev=2)
    df['BB_upper_10_2'] = bb_u10
    df['BB_middle_10_2'] = bb_m10
    df['BB_lower_10_2'] = bb_l10

    # --- Stochastic variants ---
    stoch_k5, stoch_d3 = add_stochastic(df, k_period=5, d_period=3)
    df['STOCH_K_5'] = stoch_k5
    df['STOCH_D_3'] = stoch_d3
    stoch_k8, stoch_d3v2 = add_stochastic(df, k_period=8, d_period=3)
    df['STOCH_K_8'] = stoch_k8
    df['STOCH_D_3b'] = stoch_d3v2

    # --- Candle patterns ---
    patterns = add_candle_patterns(df)
    for name, values in patterns.items():
        df[f'pattern_{name}'] = values

    # --- Candle anatomy ---
    df['candle_body'] = df['Close'] - df['Open']
    df['candle_range'] = df['High'] - df['Low']
    df['body_pct'] = np.where(df['candle_range'] > 0, abs(df['candle_body']) / df['candle_range'], 0)
    df['upper_shadow'] = df['High'] - np.maximum(df['Open'], df['Close'])
    df['lower_shadow'] = df['Low'] - np.minimum(df['Open'], df['Close'])

    # --- RSI divergence variants ---
    df = detect_rsi_divergence(df, lookback=20, out_bull='rsi_bull_div_20', out_bear='rsi_bear_div_20')
    df = detect_rsi_divergence(df, lookback=15, out_bull='rsi_bull_div_15', out_bear='rsi_bear_div_15')

    # --- S/R levels variants ---
    df = detect_sr_levels(df, lookback=20, out_res='sr_resistance_20', out_sup='sr_support_20')
    df = detect_sr_levels(df, lookback=15, out_res='sr_resistance_15', out_sup='sr_support_15')

    return df
