from __future__ import annotations

from dataclasses import dataclass

import pandas as pd


@dataclass(frozen=True)
class TimeFold:
    fold: int
    train_dates: pd.DatetimeIndex
    validation_dates: pd.DatetimeIndex
    test_dates: pd.DatetimeIndex

    def masks(self, frame: pd.DataFrame, date_column: str = "date") -> tuple[pd.Series, pd.Series, pd.Series]:
        return (
            frame[date_column].isin(self.train_dates),
            frame[date_column].isin(self.validation_dates),
            frame[date_column].isin(self.test_dates),
        )


class ExpandingWindowSplitter:
    def __init__(
        self,
        minimum_train_sessions: int,
        validation_sessions: int,
        test_sessions: int,
        step_sessions: int,
    ) -> None:
        values = (minimum_train_sessions, validation_sessions, test_sessions, step_sessions)
        if any(value <= 0 for value in values):
            raise ValueError("All split sizes must be positive")
        self.minimum_train_sessions = minimum_train_sessions
        self.validation_sessions = validation_sessions
        self.test_sessions = test_sessions
        self.step_sessions = step_sessions

    def split(self, dates: pd.Series) -> list[TimeFold]:
        sessions = pd.DatetimeIndex(pd.to_datetime(dates).dropna().unique()).sort_values()
        required = self.minimum_train_sessions + self.validation_sessions + self.test_sessions
        if len(sessions) < required:
            raise ValueError(f"Need at least {required} sessions; found {len(sessions)}")
        folds: list[TimeFold] = []
        train_end = self.minimum_train_sessions
        fold_id = 0
        while train_end + self.validation_sessions + self.test_sessions <= len(sessions):
            validation_end = train_end + self.validation_sessions
            test_end = validation_end + self.test_sessions
            folds.append(TimeFold(
                fold=fold_id,
                train_dates=sessions[:train_end],
                validation_dates=sessions[train_end:validation_end],
                test_dates=sessions[validation_end:test_end],
            ))
            fold_id += 1
            train_end += self.step_sessions
        return folds

