import matplotlib
matplotlib.use('Agg')
import mplfinance as mpf
import matplotlib.pyplot as plt
import pandas as pd
import numpy as np
from io import BytesIO


def _calc_rsi(close: pd.Series, period: int = 14) -> pd.Series:
    delta = close.diff()
    gain = delta.where(delta > 0, 0).rolling(period).mean()
    loss = (-delta.where(delta < 0, 0)).rolling(period).mean()
    rs = gain / loss
    return 100 - (100 / (1 + rs))


def render_chart(
    df: pd.DataFrame,
    title: str = "",
    buy_price: float | None = None,
    sell_price: float | None = None,
) -> BytesIO | None:
    """캔들차트 렌더링 → BytesIO PNG 반환

    df: OHLCV DataFrame (pyupbit 형식, DatetimeIndex)
    """
    if df is None or len(df) < 20:
        return None

    try:
        # pyupbit df 컬럼 확인 및 정리
        df = df.copy()
        df.index.name = "Date"
        for col in ["open", "high", "low", "close", "volume"]:
            if col not in df.columns:
                return None

        # 볼린저 밴드
        sma20 = df["close"].rolling(20).mean()
        std20 = df["close"].rolling(20).std()
        bb_upper = sma20 + 2 * std20
        bb_lower = sma20 - 2 * std20

        # RSI
        rsi = _calc_rsi(df["close"])

        # 추가 플롯
        addplots = [
            mpf.make_addplot(bb_upper, color="#888888", linestyle="--", width=0.5),
            mpf.make_addplot(bb_lower, color="#888888", linestyle="--", width=0.5),
            mpf.make_addplot(rsi, panel=2, color="#7B68EE", ylabel="RSI", width=0.8),
        ]

        # 매수가 라인
        if buy_price:
            buy_line = pd.Series([buy_price] * len(df), index=df.index)
            addplots.append(mpf.make_addplot(buy_line, color="#2196F3", linestyle=":", width=1.0))

        # 매도가 라인
        if sell_price:
            sell_line = pd.Series([sell_price] * len(df), index=df.index)
            addplots.append(mpf.make_addplot(sell_line, color="#FF5722", linestyle=":", width=1.0))

        # RSI 기준선 (30, 70)
        rsi_30 = pd.Series([30] * len(df), index=df.index)
        rsi_70 = pd.Series([70] * len(df), index=df.index)
        addplots.append(mpf.make_addplot(rsi_30, panel=2, color="#AAAAAA", linestyle="--", width=0.4))
        addplots.append(mpf.make_addplot(rsi_70, panel=2, color="#AAAAAA", linestyle="--", width=0.4))

        buf = BytesIO()
        mpf.plot(
            df,
            type="candle",
            style="charles",
            title=title,
            volume=True,
            addplot=addplots,
            figsize=(10, 7),
            panel_ratios=(4, 1, 1),
            savefig=dict(fname=buf, dpi=100, bbox_inches="tight"),
            warn_too_much_data=999,
        )
        buf.seek(0)
        plt.close("all")
        return buf

    except Exception as e:
        plt.close("all")
        return None
