"""
Модуль Dataset классов для нейронных сетей.

Классы:
    TimeSeriesDataset — стандартный датасет со скользящим окном (1 таймфрейм)
    MultiTimeframeDataset — датасет с несколькими таймфреймами (H1 + D1 + W1)
    MultiHorizonDataset — датасет с несколькими горизонтами прогноза
    WalkForwardDataset — датасет для walk-forward валидации
"""

from typing import Optional, Union

import numpy as np
import pandas as pd
import torch
from torch.utils.data import Dataset


class TimeSeriesDataset(Dataset):
    """
    PyTorch Dataset для временных рядов со скользящим окном.

    Каждый сэмпл — последовательность из seq_len наблюдений,
    за которой следует одно целевое значение.

    Args:
        features: матрица признаков (n_samples x n_features).
        targets: вектор целевых значений (n_samples,).
        seq_len: длина окна (число прошлых наблюдений).
        stride: шаг между окнами (1 = полностью перекрывающиеся).
        transform: опциональная функция трансформации признаков.

    Пример:
        >>> ds = TimeSeriesDataset(X_train, y_train, seq_len=60)
        >>> x, y = ds[0]
        >>> x.shape
        torch.Size([60, n_features])
    """

    def __init__(
        self,
        features: np.ndarray,
        targets: np.ndarray,
        seq_len: int = 60,
        stride: int = 1,
        transform: Optional[callable] = None,
    ) -> None:
        if len(features) != len(targets):
            raise ValueError(
                f"Размеры features ({len(features)}) и targets "
                f"({len(targets)}) не совпадают"
            )
        if seq_len < 1:
            raise ValueError(f"seq_len должен быть >= 1, получено {seq_len}")

        self.features = features
        self.targets = targets
        self.seq_len = seq_len
        self.stride = stride
        self.transform = transform

        # Предварительный расчёт валидных индексов
        # Индекс — это позиция ТАРГЕТА (последний элемент окна)
        valid_indices = list(range(seq_len - 1, len(features) - 1, stride))
        self.indices = valid_indices

    def __len__(self) -> int:
        """Вернуть количество сэмплов в датасете."""
        return len(self.indices)

    def __getitem__(
        self,
        idx: int,
    ) -> tuple[torch.Tensor, torch.Tensor]:
        """
        Получить сэмпл по индексу.

        Args:
            idx: индекс сэмпла.

        Returns:
            Кортеж (x, y), где:
                x: тензор признаков формы (seq_len, n_features)
                y: тензор целевого значения (скаляр)
        """
        end_idx = self.indices[idx]
        start_idx = end_idx - self.seq_len + 1

        x = self.features[start_idx:end_idx + 1]
        # Если seq_len=1, убираем лишнюю размерность
        if self.seq_len == 1:
            x = x.squeeze(0)
        y = self.targets[end_idx]

        x_tensor = torch.FloatTensor(x)
        y_tensor = torch.LongTensor([y])[0] if np.issubdtype(
            self.targets.dtype, np.integer
        ) else torch.FloatTensor([y])[0]

        if self.transform:
            x_tensor = self.transform(x_tensor)

        return x_tensor, y_tensor


class MultiTimeframeDataset(Dataset):
    """
    PyTorch Dataset для нескольких таймфреймов (H1 + D1 + W1).

    Каждый сэмпл содержит последовательности из всех таймфреймов,
    которые заканчиваются в один и тот же календарный момент.

    Args:
        h1_data: матрица признаков H1 (n_h1_samples x n_features).
        d1_data: матрица признаков D1 (n_d1_samples x n_features).
        w1_data: матрица признаков W1 (n_w1_samples x n_features).
        targets: вектор целевых значений.
        h1_seq_len: длина окна для H1 (по умолчанию 48 для 2 дней H1).
        d1_seq_len: длина окна для D1 (по умолчанию 30).
        w1_seq_len: длина окна для W1 (по умолчанию 12).
        stride: шаг между окнами.
    """

    def __init__(
        self,
        h1_data: np.ndarray,
        d1_data: np.ndarray,
        w1_data: np.ndarray,
        targets: np.ndarray,
        h1_seq_len: int = 48,
        d1_seq_len: int = 30,
        w1_seq_len: int = 12,
        stride: int = 1,
    ) -> None:
        self.h1_data = h1_data
        self.d1_data = d1_data
        self.w1_data = w1_data
        self.targets = targets
        self.h1_seq_len = h1_seq_len
        self.d1_seq_len = d1_seq_len
        self.w1_seq_len = w1_seq_len
        self.stride = stride

        # Минимальная длина истории для всех ТФ
        self.min_history = max(h1_seq_len, d1_seq_len, w1_seq_len)
        self.indices = list(
            range(self.min_history, len(targets), stride)
        )

    def __len__(self) -> int:
        """Вернуть количество сэмплов."""
        return len(self.indices)

    def __getitem__(
        self,
        idx: int,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Получить сэмпл с тремя таймфреймами.

        Returns:
            (h1_seq, d1_seq, w1_seq, target)
        """
        end_idx = self.indices[idx]

        h1_start = end_idx - self.h1_seq_len
        d1_start = end_idx - self.d1_seq_len
        w1_start = end_idx - self.w1_seq_len

        return (
            torch.FloatTensor(self.h1_data[h1_start:end_idx]),
            torch.FloatTensor(self.d1_data[d1_start:end_idx]),
            torch.FloatTensor(self.w1_data[w1_start:end_idx]),
            torch.LongTensor([self.targets[end_idx]])[0],
        )


class MultiHorizonDataset(Dataset):
    """
    PyTorch Dataset для нескольких горизонтов прогноза одновременно.

    Каждый сэмпл содержит одну последовательность признаков
    и целевые значения для нескольких горизонтов.

    Args:
        features: матрица признаков.
        targets_dict: словарь {название_горизонта: вектор_таргетов}.
        seq_len: длина окна.
        stride: шаг между окнами.
    """

    def __init__(
        self,
        features: np.ndarray,
        targets_dict: dict[str, np.ndarray],
        seq_len: int = 60,
        stride: int = 1,
    ) -> None:
        self.features = features
        self.targets_dict = targets_dict
        self.seq_len = seq_len
        self.stride = stride

        # Проверка размеров
        for name, targets in targets_dict.items():
            if len(targets) != len(features):
                raise ValueError(
                    f"Размер targets '{name}' ({len(targets)}) "
                    f"не совпадает с features ({len(features)})"
                )

        self.indices = list(range(seq_len - 1, len(features) - 1, stride))

    def __len__(self) -> int:
        return len(self.indices)

    def __getitem__(
        self,
        idx: int,
    ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
        """Получить сэмпл с несколькими горизонтами."""
        end_idx = self.indices[idx]
        start_idx = end_idx - self.seq_len + 1

        x = torch.FloatTensor(self.features[start_idx:end_idx + 1])
        y = {
            name: torch.LongTensor([targets[end_idx]])[0]
            for name, targets in self.targets_dict.items()
        }
        return x, y


class WalkForwardDataset(Dataset):
    """
    PyTorch Dataset для walk-forward валидации.

    Позволяет итерироваться по временным рядам с перекрывающимися
    окнами обучения/валидации для реалистичного бэктестинга.

    Args:
        features: полная матрица признаков.
        targets: полный вектор целей.
        seq_len: длина окна для признаков.
        train_size: размер обучающего окна (в сэмплах).
        val_size: размер валидационного окна.
        step: шаг смещения окна.
    """

    def __init__(
        self,
        features: np.ndarray,
        targets: np.ndarray,
        seq_len: int = 60,
        train_size: int = 300,
        val_size: int = 50,
        step: int = 50,
    ) -> None:
        self.features = features
        self.targets = targets
        self.seq_len = seq_len
        self.train_size = train_size
        self.val_size = val_size
        self.step = step

        total_needed = seq_len + train_size + val_size
        if len(features) < total_needed:
            raise ValueError(
                f"Недостаточно данных: нужно {total_needed}, "
                f"доступно {len(features)}"
            )

    def get_split(
        self,
        fold: int,
    ) -> tuple['TimeSeriesDataset', 'TimeSeriesDataset']:
        """
        Получить fold-разбиение для walk-forward.

        Args:
            fold: номер фолда (0, 1, 2, ...).

        Returns:
            (train_dataset, val_dataset)
        """
        start = fold * self.step
        train_end = start + self.train_size
        val_end = train_end + self.val_size

        if val_end >= len(self.features):
            raise IndexError(
                f"Фолд {fold} выходит за границы данных"
            )

        X_train = self.features[start:train_end]
        y_train = self.targets[start:train_end]
        X_val = self.features[train_end - self.seq_len + 1:val_end]
        y_val = self.targets[train_end - self.seq_len + 1:val_end]

        return (
            TimeSeriesDataset(X_train, y_train, self.seq_len),
            TimeSeriesDataset(X_val, y_val, self.seq_len),
        )

    def __len__(self) -> int:
        """Вернуть количество фолдов."""
        total = len(self.features)
        return max(0, (total - self.train_size - self.seq_len) // self.step)


def create_sequences(
    features: np.ndarray,
    targets: np.ndarray,
    seq_len: int = 60,
    stride: int = 1,
) -> tuple[np.ndarray, np.ndarray]:
    """
    Создать последовательности для скользящего окна (без PyTorch).

    Полезно для сохранения данных в HDF5 или NumPy формате.

    Args:
        features: матрица признаков (n_samples x n_features).
        targets: вектор целей (n_samples,).
        seq_len: длина окна.
        stride: шаг между окнами.

    Returns:
        Кортеж (X_seq, y_seq), где:
            X_seq: массив формы (n_sequences, seq_len, n_features)
            y_seq: массив формы (n_sequences,)
    """
    n = len(features)
    indices = list(range(seq_len - 1, n - 1, stride))

    X_list = []
    y_list = []

    for end_idx in indices:
        start_idx = end_idx - seq_len + 1
        X_list.append(features[start_idx:end_idx + 1])
        y_list.append(targets[end_idx])

    return np.array(X_list), np.array(y_list)
