"""
Модуль расчёта технических индикаторов.

Функции:
    calculate_ema — экспоненциальная скользящая средняя
    calculate_sma — простая скользящая средняя
    calculate_macd — индикатор MACD
    calculate_rsi — индекс относительной силы
    calculate_stochastic — стохастический осциллятор
    calculate_bollinger_bands — полосы Боллинджера
    calculate_atr — средний истинный диапазон
    calculate_adx — средний индекс направленности
    calculate_cci — индекс товарного канала
    calculate_obv — балансовый объём
    calculate_vwap — средневзвешенная по объёму цена
"""

from typing import Optional, Tuple

import numpy as np
import pandas as pd


def calculate_ema(close: pd.Series, period: int) -> pd.Series:
    """Расчёт экспоненциальной скользящей средней."""
    return close.ewm(span=period, adjust=False).mean()


def calculate_sma(close: pd.Series, period: int) -> pd.Series:
    """Расчёт простой скользящей средней."""
    return close.rolling(window=period).mean()


def calculate_macd(
    close: pd.Series,
    fast: int = 12,
    slow: int = 26,
    signal: int = 9,
) -> Tuple[pd.Series, pd.Series, pd.Series]:
    """
    Расчёт индикатора MACD.

    Returns:
        (macd_line, signal_line, histogram)
    """
    ema_fast = calculate_ema(close, fast)
    ema_slow = calculate_ema(close, slow)
    macd_line = ema_fast - ema_slow
    signal_line = calculate_ema(macd_line, signal)
    histogram = macd_line - signal_line
    return macd_line, signal_line, histogram


def calculate_rsi(close: pd.Series, period: int = 14) -> pd.Series:
    """Расчёт индикатора RSI (0-100)."""
    delta = close.diff()
    gain = delta.clip(lower=0)
    loss = -delta.clip(upper=0)
    avg_gain = gain.ewm(alpha=1 / period, adjust=False).mean()
    avg_loss = loss.ewm(alpha=1 / period, adjust=False).mean()
    rs = avg_gain / avg_loss
    return 100 - (100 / (1 + rs))


def calculate_stochastic(
    high: pd.Series,
    low: pd.Series,
    close: pd.Series,
    k_period: int = 14,
    d_period: int = 3,
) -> Tuple[pd.Series, pd.Series]:
    """
    Расчёт стохастического осциллятора.

    Returns:
        (stoch_k, stoch_d)
    """
    lowest_low = low.rolling(window=k_period).min()
    highest_high = high.rolling(window=k_period).max()
    stoch_k = 100 * (close - lowest_low) / (highest_high - lowest_low)
    stoch_d = stoch_k.rolling(window=d_period).mean()
    return stoch_k, stoch_d


def calculate_bollinger_bands(
    close: pd.Series,
    period: int = 20,
    num_std: int = 2,
) -> Tuple[pd.Series, pd.Series, pd.Series]:
    """
    Расчёт полос Боллинджера.

    Returns:
        (middle_band, upper_band, lower_band)
    """
    middle = calculate_sma(close, period)
    std = close.rolling(window=period).std()
    upper = middle + num_std * std
    lower = middle - num_std * std
    return middle, upper, lower


def calculate_atr(
    high: pd.Series,
    low: pd.Series,
    close: pd.Series,
    period: int = 14,
) -> pd.Series:
    """Расчёт среднего истинного диапазона (ATR)."""
    prev_close = close.shift(1)
    tr1 = high - low
    tr2 = (high - prev_close).abs()
    tr3 = (low - prev_close).abs()
    true_range = pd.concat([tr1, tr2, tr3], axis=1).max(axis=1)
    return true_range.ewm(alpha=1 / period, adjust=False).mean()


def calculate_adx(
    high: pd.Series,
    low: pd.Series,
    close: pd.Series,
    period: int = 14,
) -> Tuple[pd.Series, pd.Series, pd.Series]:
    """
    Расчёт среднего индекса направленности.

    Returns:
        (adx, plus_di, minus_di)
    """
    prev_high = high.shift(1)
    prev_low = low.shift(1)
    prev_close = close.shift(1)

    plus_dm = np.where((high - prev_high) > (prev_low - low), np.maximum(high - prev_high, 0), 0)
    minus_dm = np.where((prev_low - low) > (high - prev_high), np.maximum(prev_low - low, 0), 0)

    tr1 = high - low
    tr2 = (high - prev_close).abs()
    tr3 = (low - prev_close).abs()
    true_range = pd.concat([pd.Series(tr1), pd.Series(tr2), pd.Series(tr3)], axis=1).max(axis=1)

    atr_val = true_range.ewm(alpha=1 / period, adjust=False).mean()
    plus_di = pd.Series(100 * pd.Series(plus_dm).ewm(alpha=1 / period, adjust=False).mean() / atr_val)
    minus_di = pd.Series(100 * pd.Series(minus_dm).ewm(alpha=1 / period, adjust=False).mean() / atr_val)

    dx = 100 * (plus_di - minus_di).abs() / (plus_di + minus_di)
    adx = dx.ewm(alpha=1 / period, adjust=False).mean()

    return adx, plus_di, minus_di


def calculate_cci(
    high: pd.Series,
    low: pd.Series,
    close: pd.Series,
    period: int = 20,
) -> pd.Series:
    """Расчёт индекса товарного канала (CCI)."""
    typical_price = (high + low + close) / 3
    sma_tp = calculate_sma(typical_price, period)
    mean_deviation = (typical_price - sma_tp).abs().rolling(window=period).mean()
    return (typical_price - sma_tp) / (0.015 * mean_deviation)


def calculate_obv(close: pd.Series, volume: pd.Series) -> pd.Series:
    """Расчёт балансового объёма (OBV)."""
    direction = np.where(close.diff() > 0, 1, np.where(close.diff() < 0, -1, 0))
    obv = (direction * volume).cumsum()
    return pd.Series(obv, index=close.index)


def detect_swing_levels(
    high: pd.Series,
    low: pd.Series,
    window: int = 5,
) -> Tuple[pd.Series, pd.Series]:
    """
    Поиск свинг-уровней (экстремумов).

    Returns:
        (swing_highs, swing_lows)
    """
    swing_highs = pd.Series(np.nan, index=high.index)
    swing_lows = pd.Series(np.nan, index=low.index)

    for i in range(window, len(high) - window):
        if high.iloc[i] == high.iloc[i - window : i + window + 1].max():
            swing_highs.iloc[i] = high.iloc[i]
        if low.iloc[i] == low.iloc[i - window : i + window + 1].min():
            swing_lows.iloc[i] = low.iloc[i]

    return swing_highs, swing_lows


def calculate_volume_sma(volume: pd.Series, period: int = 20) -> pd.Series:
    """Расчёт скользящей средней объёма."""
    return calculate_sma(volume, period)


def calc_all_indicators(df: pd.DataFrame) -> pd.DataFrame:
    """
    Расчёт полного комплекта индикаторов для DataFrame с OHLCV.

    Args:
        df: DataFrame с колонками [Open, High, Low, Close, Volume].

    Returns:
        DataFrame с добавленными колонками индикаторов.
    """
    df = df.copy()

    # Трендовые
    for period in [9, 21, 50, 200]:
        if len(df) >= period:
            df[f'EMA_{period}'] = calculate_ema(df['Close'], period)

    for period in [20, 50, 200]:
        if len(df) >= period:
            df[f'SMA_{period}'] = calculate_sma(df['Close'], period)

    # MACD
    if len(df) >= 26:
        df['MACD_line'], df['MACD_signal'], df['MACD_hist'] = calculate_macd(df['Close'])

    # RSI
    if len(df) >= 14:
        df['RSI_14'] = calculate_rsi(df['Close'], 14)

    # Stochastic
    if len(df) >= 14:
        df['Stoch_K'], df['Stoch_D'] = calculate_stochastic(df['High'], df['Low'], df['Close'])

    # Bollinger Bands
    if len(df) >= 20:
        _, df['BB_upper'], df['BB_lower'] = calculate_bollinger_bands(df['Close'])

    # ATR
    if len(df) >= 14:
        df['ATR_14'] = calculate_atr(df['High'], df['Low'], df['Close'], 14)

    # ADX
    if len(df) >= 14:
        df['ADX_14'], df['Plus_DI'], df['Minus_DI'] = calculate_adx(df['High'], df['Low'], df['Close'])

    # CCI
    if len(df) >= 20:
        df['CCI_20'] = calculate_cci(df['High'], df['Low'], df['Close'])

    # OBV
    df['OBV'] = calculate_obv(df['Close'], df['Volume'])

    # Объём
    df['Vol_SMA_20'] = calculate_volume_sma(df['Volume'], 20)

    return df
