"""퀀티랩 기술적 분석 시리즈 공개 재현 코드.

원본 OHLCV CSV를 읽어 이동평균, 볼린저 밴드, MACD, RSI, ATR, OBV와
간단한 이벤트 수익률을 계산합니다.
"""

from pathlib import Path

import numpy as np
import pandas as pd


def add_indicators(df: pd.DataFrame) -> pd.DataFrame:
    df = df.sort_values(["code", "date"]).copy()
    close = df["close"].astype(float)
    high = df["high"].astype(float)
    low = df["low"].astype(float)
    volume = df["volume"].astype(float)

    for window in (5, 20, 60, 120):
        df[f"ma{window}"] = close.groupby(df["code"]).transform(
            lambda x: x.rolling(window).mean()
        )

    def per_stock(group: pd.DataFrame) -> pd.DataFrame:
        group = group.sort_values("date").copy()
        close = group["close"]
        high = group["high"]
        low = group["low"]
        volume = group["volume"]

        middle = close.rolling(20).mean()
        std = close.rolling(20).std()
        group["bb_middle"] = middle
        group["bb_upper"] = middle + 2 * std
        group["bb_lower"] = middle - 2 * std

        fast = close.ewm(span=12).mean()
        slow = close.ewm(span=26).mean()
        group["macd"] = fast - slow
        group["macd_signal"] = group["macd"].ewm(span=9).mean()

        change = close.diff()
        gain = change.clip(lower=0).rolling(14).mean()
        loss = (-change.clip(upper=0)).rolling(14).mean()
        group["rsi14"] = gain / (gain + loss) * 100

        true_range = pd.concat(
            [high - low, (high - close.shift()).abs(), (low - close.shift()).abs()],
            axis=1,
        ).max(axis=1)
        group["atr14_pct"] = true_range.rolling(14).mean() / close
        group["obv"] = (np.sign(close.diff()).fillna(0) * volume).cumsum()
        group["obv_ma20"] = group["obv"].rolling(20).mean()
        group["volume_ratio20"] = volume / volume.rolling(20).mean()
        group["return_1d"] = close.pct_change()

        group["ma_trend"] = ((close > group["ma20"]) & (group["ma20"] > group["ma60"])).astype(int)
        group["rsi_recovery"] = (
            (group["rsi14"].shift(1) <= 30) & (group["rsi14"] > 30)
        ).astype(int)
        group["macd_golden_cross"] = (
            (group["macd"].shift(1) <= group["macd_signal"].shift(1))
            & (group["macd"] > group["macd_signal"])
        ).astype(int)
        group["volume_breakout"] = (
            (close > close.rolling(20).max().shift(1))
            & (volume > volume.rolling(20).mean().shift(1) * 1.5)
        ).astype(int)
        group["composite_state"] = (
            (close > group["ma20"])
            & (group["ma20"] > group["ma60"])
            & group["rsi14"].between(40, 70)
            & (group["macd"] > group["macd_signal"])
            & (group["volume_ratio20"] > 1)
        ).astype(int)
        for horizon in (5, 20, 60):
            group[f"forward_return_{horizon}d"] = close.shift(-horizon) / close - 1
        return group

    return df.groupby("code", group_keys=False).apply(per_stock, include_groups=True)


def summarize_events(df: pd.DataFrame, signal: str, horizon: int = 20) -> pd.Series:
    values = df.loc[df[signal].eq(1), f"forward_return_{horizon}d"].dropna()
    return pd.Series(
        {
            "events": len(values),
            "positive_rate": values.gt(0).mean(),
            "mean_return": values.mean(),
            "median_return": values.median(),
        }
    )


if __name__ == "__main__":
    root = Path(__file__).resolve().parent
    raw = pd.read_csv(root / "ohlcv.csv", dtype={"code": str})
    result = add_indicators(raw)
    print(summarize_events(result, "macd_golden_cross").to_string())
