# 來源註記：這是 Codex 弄的。
"""教學用績效指標；輸入須為有 DatetimeIndex 的 equity curve。"""

from __future__ import annotations

import math

import numpy as np
import pandas as pd


def max_drawdown_duration(drawdown: pd.Series) -> int:
    """回傳連續處於水下的最長觀測期數。"""
    longest = current = 0
    for value in drawdown.fillna(0.0):
        if value < 0:
            current += 1
            longest = max(longest, current)
        else:
            current = 0
    return longest


def performance_metrics(
    equity: pd.Series,
    periods_per_year: int = 252,
    annual_risk_free_rate: float = 0.0,
) -> dict[str, float]:
    equity = equity.dropna().astype(float).sort_index()
    if len(equity) < 2 or (equity <= 0).any():
        raise ValueError("equity 至少需兩筆且全部大於 0")
    if not isinstance(equity.index, pd.DatetimeIndex):
        raise TypeError("equity.index 必須是 DatetimeIndex")

    returns = equity.pct_change().dropna()
    elapsed_years = (equity.index[-1] - equity.index[0]).total_seconds() / (
        365.2425 * 24 * 60 * 60
    )
    if elapsed_years <= 0:
        raise ValueError("時間範圍必須大於 0")

    cagr = (equity.iloc[-1] / equity.iloc[0]) ** (1 / elapsed_years) - 1
    period_rf = (1 + annual_risk_free_rate) ** (1 / periods_per_year) - 1
    excess = returns - period_rf
    period_std = returns.std(ddof=1)
    annual_vol = period_std * math.sqrt(periods_per_year)
    sharpe = (
        excess.mean() / period_std * math.sqrt(periods_per_year)
        if period_std > 0
        else math.nan
    )

    downside = np.minimum(excess.to_numpy(), 0.0)
    annual_downside = np.sqrt(np.mean(downside**2)) * math.sqrt(periods_per_year)
    sortino = (
        excess.mean() * periods_per_year / annual_downside
        if annual_downside > 0
        else math.nan
    )

    running_peak = equity.cummax()
    drawdown = equity / running_peak - 1
    max_dd = float(drawdown.min())
    calmar = cagr / abs(max_dd) if max_dd < 0 else math.nan

    return {
        "CAGR": float(cagr),
        "Annual volatility": float(annual_vol),
        "Sharpe": float(sharpe),
        "Sortino": float(sortino),
        "Max drawdown": max_dd,
        "Max drawdown duration (periods)": float(max_drawdown_duration(drawdown)),
        "Calmar": float(calmar),
    }


def _demo() -> None:
    rng = np.random.default_rng(42)
    dates = pd.bdate_range("2022-01-03", periods=756)
    returns = rng.normal(0.00035, 0.012, len(dates))
    equity = pd.Series(100_000 * np.cumprod(1 + returns), index=dates)
    metrics = performance_metrics(equity, annual_risk_free_rate=0.03)
    for name, value in metrics.items():
        print(f"{name:36s} {value: .4f}")


if __name__ == "__main__":
    _demo()

