# 來源註記：這是 Codex 弄的。
"""教學用 rolling walk-forward；每個 test fold 只使用過去資料選參數。"""

from __future__ import annotations

import math

import numpy as np
import pandas as pd


def strategy_returns(close: pd.Series, window: int, cost_bps: float = 5.0) -> pd.Series:
    signal = (close > close.rolling(window, min_periods=window).mean()).astype(float)
    position = signal.shift(1).fillna(0.0)
    turnover = position.diff().abs().fillna(position.abs())
    return position * close.pct_change().fillna(0.0) - turnover * cost_bps / 10_000


def annualized_sharpe(returns: pd.Series, periods_per_year: int = 252) -> float:
    returns = returns.dropna()
    std = returns.std(ddof=1)
    return returns.mean() / std * math.sqrt(periods_per_year) if std > 0 else -math.inf


def rolling_walk_forward(
    close: pd.Series,
    windows: tuple[int, ...] = (10, 20, 40, 80),
    train_size: int = 252,
    test_size: int = 63,
) -> tuple[pd.Series, pd.DataFrame]:
    folds: list[dict[str, object]] = []
    oos_parts: list[pd.Series] = []
    start = train_size
    while start + test_size <= len(close):
        train = close.iloc[start - train_size : start]
        # 在 train 內選參數；完整 rolling warm-up 不可偷用 test。
        scores = {window: annualized_sharpe(strategy_returns(train, window)) for window in windows}
        best_window = max(scores, key=scores.get)

        warmup_start = max(0, start - best_window)
        combined = close.iloc[warmup_start : start + test_size]
        combined_returns = strategy_returns(combined, best_window)
        test_index = close.iloc[start : start + test_size].index
        test_returns = combined_returns.reindex(test_index)
        oos_parts.append(test_returns)
        folds.append(
            {
                "train_start": train.index[0],
                "train_end": train.index[-1],
                "test_start": test_index[0],
                "test_end": test_index[-1],
                "best_window": best_window,
                "train_sharpe": scores[best_window],
                "test_sharpe": annualized_sharpe(test_returns),
            }
        )
        start += test_size
    if not oos_parts:
        raise ValueError("資料長度不足以建立一個 fold")
    return pd.concat(oos_parts).sort_index(), pd.DataFrame(folds)


def _demo() -> None:
    rng = np.random.default_rng(123)
    dates = pd.bdate_range("2018-01-02", periods=1_260)
    drift = np.select(
        [np.arange(len(dates)) < 420, np.arange(len(dates)) < 840],
        [0.0005, -0.0001],
        default=0.00025,
    )
    close = pd.Series(100 * np.cumprod(1 + drift + rng.normal(0, 0.012, len(dates))), index=dates)
    oos, folds = rolling_walk_forward(close)
    print(folds.to_string(index=False))
    print(f"Combined OOS Sharpe: {annualized_sharpe(oos):.3f}")
    assert oos.index.is_monotonic_increasing and not oos.index.has_duplicates
    assert folds["test_start"].gt(folds["train_end"]).all()


if __name__ == "__main__":
    _demo()

