"""
Cross-Validation для временных рядов.

Предотвращает Data Leakage при кросс-валидации, используя скользящее окно с зазором.
"""

import numpy as np
from typing import Iterator, Tuple
from utils.logger import logger


class WalkForwardSplit:
    """
    Walk-Forward Cross-Validation для временных рядов.
    Обучается на расширяющемся (или скользящем) окне, тестируется на следующем.
    Поддерживает 'gap' для предотвращения Data Leakage от будущих таргетов.

    Args:
        n_splits: Количество фолдов
        test_size: Размер тестового окна (float или int)
        gap: Зазор между train и test в баров

    Example:
        >>> cv = WalkForwardSplit(n_splits=5, test_size=100, gap=5)
        >>> for train_idx, test_idx in cv.split(X):
        ...     X_train, X_test = X[train_idx], X[test_idx]
        ...     model.fit(X_train, y_train)
        ...     y_pred = model.predict(X_test)
    """

    def __init__(self, n_splits: int = 5, test_size: float = 0.2, gap: int = 5):
        """
        Инициализация Walk-Forward сплита.

        Args:
            n_splits: Количество фолдов
            test_size: Размер тестового окна (0.2 = 20% от данных)
            gap: Зазор между train и test (баров)
        """
        self.n_splits = n_splits
        self.test_size = test_size
        self.gap = gap

    def split(self, X, y=None, groups=None) -> Iterator[Tuple[np.ndarray, np.ndarray]]:
        """
        Генератор фолдов для кросс-валидации.

        Args:
            X: Feature matrix (n_samples, n_features)
            y: Target vector or matrix
            groups: Not used (для совместимости с sklearn)

        Yields:
            Tuple of (train_indices, test_indices)
        """
        n_samples = len(X)

        # Определяем размер тестового окна
        if isinstance(self.test_size, float):
            test_size = int(n_samples * self.test_size)
        else:
            test_size = self.test_size

        # Размер обучающего окна для первого фолда
        # Оставляем достаточно данных для test_size * n_splits + gap * n_splits
        initial_train_size = n_samples - (test_size * self.n_splits) - (self.gap * self.n_splits)

        if initial_train_size < test_size:
            raise ValueError(
                f"Недостаточно данных для заданного n_splits={self.n_splits}, "
                f"test_size={test_size}, gap={self.gap}. "
                f"Нужно минимум {test_size + self.gap} баров."
            )

        logger.info(f"Walk-Forward Split: n_splits={self.n_splits}, test_size={test_size}, gap={self.gap}")

        for i in range(self.n_splits):
            train_end = initial_train_size + i * test_size
            test_start = train_end + self.gap
            test_end = test_start + test_size

            if test_end > n_samples:
                test_end = n_samples

            train_indices = np.arange(0, train_end)
            test_indices = np.arange(test_start, test_end)

            logger.debug(f"Fold {i+1}: Train={len(train_indices)}, Test={len(test_indices)}")

            yield train_indices, test_indices

    def get_n_splits(self):
        """Возвращает количество фолдов."""
        return self.n_splits

    def get_params(self):
        """Возвращает параметры."""
        return {
            'n_splits': self.n_splits,
            'test_size': self.test_size,
            'gap': self.gap
        }

    def __repr__(self):
        return f"WalkForwardSplit(n_splits={self.n_splits}, test_size={self.test_size}, gap={self.gap})"
