"""技术指标，纯 pandas 实现（避免 talib 编译依赖）。

与 freqtrade/talib 口径对齐：
- ATR: Wilder 平滑（alpha=1/n）
- ADX: Wilder 平滑
- EMA/SMA/Donchian: 标准定义
"""
from __future__ import annotations

import numpy as np
import pandas as pd


def ema(close: pd.Series, period: int) -> pd.Series:
    return close.ewm(span=period, adjust=False).mean()


def sma(close: pd.Series, period: int) -> pd.Series:
    return close.rolling(period).mean()


def true_range(df: pd.DataFrame) -> pd.Series:
    high_low = df["high"] - df["low"]
    high_close = (df["high"] - df["close"].shift(1)).abs()
    low_close = (df["low"] - df["close"].shift(1)).abs()
    return pd.concat([high_low, high_close, low_close], axis=1).max(axis=1)


def atr(df: pd.DataFrame, period: int = 14) -> pd.Series:
    """Wilder ATR，与 talib 一致。"""
    tr = true_range(df)
    return tr.ewm(alpha=1 / period, min_periods=period, adjust=False).mean()


def adx(df: pd.DataFrame, period: int = 14) -> pd.Series:
    """Wilder ADX，与 talib 一致。"""
    up = df["high"].diff()
    down = -df["low"].diff()
    plus_dm = np.where((up > down) & (up > 0), up, 0.0)
    minus_dm = np.where((down > up) & (down > 0), down, 0.0)
    tr = true_range(df)
    atr_ = tr.ewm(alpha=1 / period, min_periods=period, adjust=False).mean()
    plus_di = 100 * pd.Series(plus_dm, index=df.index).ewm(
        alpha=1 / period, min_periods=period, adjust=False).mean() / atr_
    minus_di = 100 * pd.Series(minus_dm, index=df.index).ewm(
        alpha=1 / period, min_periods=period, adjust=False).mean() / atr_
    dx = 100 * (plus_di - minus_di).abs() / (plus_di + minus_di).replace(0, np.nan)
    return dx.ewm(alpha=1 / period, min_periods=period, adjust=False).mean()


def donchian_high(df: pd.DataFrame, period: int) -> pd.Series:
    """前 N 根最高价，shift(1) 排除当前根，避免未来数据。"""
    return df["high"].rolling(period).max().shift(1)


def donchian_low(df: pd.DataFrame, period: int) -> pd.Series:
    return df["low"].rolling(period).min().shift(1)


def rolling_range_pct(df: pd.DataFrame, period: int) -> pd.Series:
    """滚动窗口 (max(high)-min(low))/close，用于网格区间适宜性判断。"""
    return (df["high"].rolling(period).max() - df["low"].rolling(period).min()) / df["close"]
