"""
Аугментация временных рядов для расширения обучающей выборки.

Подходы:
    GaussianNoise        — добавление гауссова шума (имитация микроструктуры)
    TimeWarping          — временная деформация (растяжение/сжатие во времени)
    MagnitudeWarping     — деформация амплитуды (изменение силы сигнала)
    ScalingAugmentation  — масштабирование всей последовательности
    MixupAugmentation    — смешивание двух последовательностей (Mixup)
    WindowDropout        — случайное зануление участка окна
    FlipAugmentation     — инвертирование последовательности (для симметричных паттернов)

Пример использования:
    augmenter = TimeSeriesAugmenter(
        noise_std=0.005,
        warping_sigma=0.2,
        scaling_range=(0.95, 1.05),
        mixup_alpha=0.2,
    )
    X_aug, y_aug = augmenter.augment(X, y, methods=['noise', 'warp', 'mixup'])
"""

from typing import Optional, Union

import numpy as np
import torch
from scipy.interpolate import CubicSpline


class GaussianNoise:
    """
    Добавление гауссова шума к временным рядам.

    Имитирует микроструктурный шум рынка — случайные колебания цены,
    которые не несут сигнала. Увеличивает робастность модели.

    Args:
        std: стандартное отклонение шума (доля от std признака).
             Если None — вычисляется автоматически как 1% от std каждого признака.
    """

    def __init__(self, std: Optional[float] = None):
        self.std = std

    def __call__(self, X: np.ndarray, n_copies: int = 1) -> np.ndarray:
        """
        Добавить шум к последовательностям.

        Args:
            X: входные данные (N, seq_len, n_features).
            n_copies: количество копий с разным шумом.

        Returns:
            Массив (N * n_copies, seq_len, n_features).
        """
        copies = [X]
        for _ in range(n_copies):
            noise_std = self.std
            if noise_std is None:
                # Автоматический расчёт: 0.5% от std каждого признака
                feat_std = np.nanstd(X, axis=(0, 1), keepdims=True) + 1e-8
                noise_std = 0.005 * feat_std
            noise = np.random.normal(0, noise_std, size=X.shape).astype(X.dtype)
            copies.append(X + noise)
        return np.concatenate(copies, axis=0)


class TimeWarping:
    """
    Временная деформация (time warping) для временных рядов.

    Растягивает и сжимает отдельные участки временного ряда,
    имитируя разную скорость развития паттернов.
    Использует интерполяцию по ключевым точкам (knot points).

    Args:
        sigma: степень деформации (0.1-0.3).
        knot_points: количество узлов интерполяции (4-8).
    """

    def __init__(self, sigma: float = 0.2, knot_points: int = 5):
        self.sigma = sigma
        self.knot_points = knot_points

    def __call__(self, X: np.ndarray, n_copies: int = 1) -> np.ndarray:
        """
        Применить временную деформацию.

        Args:
            X: (N, seq_len, n_features).
            n_copies: количество копий.

        Returns:
            (N * (1 + n_copies), seq_len, n_features).
        """
        copies = [X]
        N, seq_len, n_features = X.shape

        for _ in range(n_copies):
            X_warped = np.zeros_like(X)
            for i in range(N):
                # Генерируем узлы деформации (индексы ключевых точек)
                knot_indices = np.linspace(0, seq_len - 1, self.knot_points).astype(int)
                orig_nodes = np.arange(seq_len, dtype=float)

                # Значения в ключевых точках
                knot_values = np.zeros((self.knot_points, n_features))
                for f in range(n_features):
                    knot_values[:, f] = X[i, knot_indices, f]

                # Деформируем положение ключевых точек (но не их значения)
                warp_positions = knot_indices.astype(float).copy()
                if self.knot_points > 2:
                    # Деформируем внутренние узлы
                    for j in range(1, self.knot_points - 1):
                        warp_positions[j] += np.random.normal(
                            0, self.sigma * seq_len / self.knot_points
                        )
                    # Проверяем монотонность
                    warp_positions = np.sort(warp_positions)
                    warp_positions[0] = 0
                    warp_positions[-1] = seq_len - 1

                # Интерполируем весь ряд через ключевые точки
                for f in range(n_features):
                    cs = CubicSpline(warp_positions, knot_values[:, f], bc_type='natural')
                    X_warped[i, :, f] = cs(orig_nodes)

            copies.append(X_warped)

        return np.concatenate(copies, axis=0)


class ScalingAugmentation:
    """
    Масштабирование амплитуды временного ряда.

    Умножает все значения на случайный коэффициент,
    имитируя разную волатильность.

    Args:
        scale_range: диапазон коэффициентов масштабирования.
    """

    def __init__(self, scale_range: tuple = (0.95, 1.05)):
        self.scale_range = scale_range

    def __call__(self, X: np.ndarray, n_copies: int = 1) -> np.ndarray:
        copies = [X]
        N, seq_len, n_features = X.shape

        for _ in range(n_copies):
            # Разный коэффициент для каждого сэмпла и признака
            scales = np.random.uniform(
                self.scale_range[0],
                self.scale_range[1],
                size=(N, 1, n_features),
            ).astype(X.dtype)
            copies.append(X * scales)

        return np.concatenate(copies, axis=0)


class MixupAugmentation:
    """
    Mixup для временных рядов.

    Линейная интерполяция между двумя случайными последовательностями.
    Целевые метки также интерполируются (soft labels).

    Args:
        alpha: параметр Beta-распределения (0.1-0.4).
    """

    def __init__(self, alpha: float = 0.2):
        self.alpha = alpha

    def __call__(
        self,
        X: np.ndarray,
        y: np.ndarray,
        n_additional: Optional[int] = None,
    ) -> tuple[np.ndarray, np.ndarray]:
        """
        Применить Mixup к батчу.

        Если n_additional указан — генерируется указанное количество
        mixup-сэмплов. Иначе — столько же, сколько в X.

        Args:
            X: (N, seq_len, n_features).
            y: (N,) — классы (0, 1, 2).
            n_additional: количество mixup-сэмплов.

        Returns:
            X_mix, y_mix_soft — объединённые с оригинальными.
        """
        N = X.shape[0]
        n_mix = n_additional if n_additional is not None else N

        # One-hot encoding для целевых меток
        y_onehot = np.zeros((N, 3), dtype=np.float32)
        y_onehot[np.arange(N), y] = 1.0

        X_mix_list = [X]
        y_mix_list = [y_onehot]

        for _ in range(n_mix // N + 1):
            # Случайные индексы
            idx = np.random.permutation(N)
            # Коэффициент смешивания
            lam = np.random.beta(self.alpha, self.alpha, size=(N, 1, 1)).astype(X.dtype)

            X_mixed = lam * X + (1 - lam) * X[idx]
            # lam: (N,1,1) → reshape to (N,1) for broadcasting with (N,3)
            lam_2d = lam.reshape(-1, 1)  # (N, 1)
            y_mixed = lam_2d * y_onehot + (1 - lam_2d) * y_onehot[idx]

            X_mix_list.append(X_mixed)
            y_mix_list.append(y_mixed)

            if sum(len(x) for x in X_mix_list) >= N + (n_additional or N):
                break

        X_result = np.concatenate(X_mix_list, axis=0)[:N + (n_additional or N)]
        y_result = np.concatenate(y_mix_list, axis=0)[:N + (n_additional or N)]

        return X_result, y_result


class WindowDropout:
    """
    Случайное зануление участка окна.

    Имитирует пропуски данных и учит модель работать с неполной информацией.

    Args:
        dropout_p: вероятность зануления признака (0.05-0.15).
        max_width: максимальная ширина зануляемого участка (5-15).
    """

    def __init__(self, dropout_p: float = 0.1, max_width: int = 10):
        self.dropout_p = dropout_p
        self.max_width = max_width

    def __call__(self, X: np.ndarray, n_copies: int = 1) -> np.ndarray:
        copies = [X]
        N, seq_len, n_features = X.shape

        for _ in range(n_copies):
            X_drop = X.copy()
            for i in range(N):
                for f in range(n_features):
                    if np.random.random() < self.dropout_p:
                        width = np.random.randint(1, self.max_width + 1)
                        start = np.random.randint(0, seq_len - width + 1)
                        X_drop[i, start:start + width, f] = 0.0
            copies.append(X_drop)

        return np.concatenate(copies, axis=0)


class TimeSeriesAugmenter:
    """
    Комбинированный аугментатор временных рядов.

    Применяет несколько методов аугментации последовательно.

    Args:
        noise_std: std для GaussianNoise (None = авто).
        warping_sigma: σ для TimeWarping.
        scaling_range: диапазон для ScalingAugmentation.
        mixup_alpha: α для Mixup.
        dropout_p: вероятность WindowDropout.
    """

    def __init__(
        self,
        noise_std: Optional[float] = None,
        warping_sigma: float = 0.2,
        scaling_range: tuple = (0.95, 1.05),
        mixup_alpha: float = 0.2,
        dropout_p: float = 0.1,
    ):
        self.noise = GaussianNoise(std=noise_std)
        self.warp = TimeWarping(sigma=warping_sigma)
        self.scale = ScalingAugmentation(scale_range=scaling_range)
        self.mixup = MixupAugmentation(alpha=mixup_alpha)
        self.dropout = WindowDropout(dropout_p=dropout_p)

    def augment(
        self,
        X: np.ndarray,
        y: np.ndarray,
        methods: Optional[list[str]] = None,
    ) -> tuple[np.ndarray, np.ndarray]:
        """
        Применить аугментацию.

        Args:
            X: (N, seq_len, n_features).
            y: (N,) — классы.
            methods: список методов ('noise', 'warp', 'scale', 'mixup', 'dropout').
                     Если None — применяются все, кроме mixup (требует y).

        Returns:
            X_aug, y_aug — расширенный датасет.
        """
        if methods is None:
            methods = ['noise', 'warp', 'scale', 'dropout']

        X_all = [X]
        y_all = [y]

        for method in methods:
            if method == 'noise':
                X_all.append(self.noise(X, n_copies=1)[len(X):])
                y_all.append(y)
            elif method == 'warp':
                X_all.append(self.warp(X, n_copies=1)[len(X):])
                y_all.append(y)
            elif method == 'scale':
                X_all.append(self.scale(X, n_copies=1)[len(X):])
                y_all.append(y)
            elif method == 'dropout':
                X_all.append(self.dropout(X, n_copies=1)[len(X):])
                y_all.append(y)
            elif method == 'mixup':
                X_mix, y_mix_soft = self.mixup(X, y, n_additional=len(X))
                X_all.append(X_mix[len(X):])
                # Для mixup конвертируем soft labels обратно в hard
                y_mix_hard = np.argmax(y_mix_soft[len(y):], axis=1)
                y_all.append(y_mix_hard)

        return np.concatenate(X_all, axis=0), np.concatenate(y_all, axis=0)

    def augment_dataloader(
        self,
        dataset: torch.utils.data.Dataset,
        batch_size: int = 32,
        factor: int = 3,
    ) -> torch.utils.data.DataLoader:
        """
        Создать DataLoader с аугментированными данными.

        Args:
            dataset: исходный TimeSeriesDataset.
            batch_size: размер батча.
            factor: коэффициент расширения.

        Returns:
            DataLoader с аугментированными данными.
        """
        # Собираем все данные из датасета
        all_X, all_y = [], []
        loader = torch.utils.data.DataLoader(dataset, batch_size=len(dataset))
        for X_batch, y_batch in loader:
            all_X.append(X_batch.numpy())
            all_y.append(y_batch.numpy())

        X = np.concatenate(all_X, axis=0)
        y = np.concatenate(all_y, axis=0)

        # Аугментируем
        X_aug, y_aug = self.augment(X, y)

        # Ограничиваем размер
        max_samples = len(X) * factor
        if len(X_aug) > max_samples:
            idx = np.random.choice(len(X_aug), max_samples, replace=False)
            X_aug, y_aug = X_aug[idx], y_aug[idx]

        # Создаём новый датасет
        aug_dataset = torch.utils.data.TensorDataset(
            torch.FloatTensor(X_aug),
            torch.LongTensor(y_aug),
        )

        return torch.utils.data.DataLoader(
            aug_dataset, batch_size=batch_size, shuffle=True
        )


def augment_dataset_hdf5(
    h5_path: str,
    output_path: str,
    methods: Optional[list[str]] = None,
    factor: int = 3,
) -> dict:
    """
    Загрузить датасет из HDF5, аугментировать и сохранить.

    Args:
        h5_path: путь к исходному HDF5.
        output_path: путь для сохранения аугментированного датасета.
        methods: методы аугментации.
        factor: коэффициент расширения.

    Returns:
        Метаданные результата.
    """
    import h5py

    augmenter = TimeSeriesAugmenter()

    with h5py.File(h5_path, 'r') as f:
        X_train = f['X_train'][:]
        y_train = f['y_train'][:]
        X_val = f['X_val'][:]
        y_val = f['y_val'][:]
        X_test = f['X_test'][:]
        y_test = f['y_test'][:]
        feature_names = f['feature_names'][:]
        scaler_mean = f['scaler_mean'][:]
        scaler_scale = f['scaler_scale'][:]
        metadata = json.loads(f['metadata'][()])

    # Аугментируем только train
    print(f"Аугментация train: {len(X_train)} → ", end="", flush=True)
    X_train_aug, y_train_aug = augmenter.augment(X_train, y_train, methods=methods)

    # Ограничиваем
    if len(X_train_aug) > len(X_train) * factor:
        idx = np.random.choice(len(X_train_aug), len(X_train) * factor, replace=False)
        X_train_aug, y_train_aug = X_train_aug[idx], y_train_aug[idx]
    print(f"{len(X_train_aug)} (+{len(X_train_aug) - len(X_train)})")

    # Сохраняем
    metadata['augmentation'] = {
        'methods': methods,
        'factor': factor,
        'original_train': int(len(X_train)),
        'augmented_train': int(len(X_train_aug)),
    }

    from src.ml.features.pipeline import save_dataset_to_hdf5
    save_dataset_to_hdf5(
        save_path=output_path,
        X_train=X_train_aug, y_train=y_train_aug,
        X_val=X_val, y_val=y_val,
        X_test=X_test, y_test=y_test,
        feature_names=list(feature_names),
        scaler_mean=scaler_mean,
        scaler_scale=scaler_scale,
        metadata=metadata,
    )

    return metadata


if __name__ == '__main__':
    import json

    # Демонстрация аугментации на X5 D1 датасете
    h5_path = '/home/ai/projects/Market Analisys/src/ml/models/saved/x5_d1_dataset.h5'
    import os
    if os.path.exists(h5_path):
        import h5py
        with h5py.File(h5_path, 'r') as f:
            X_train = f['X_train'][:]
            y_train = f['y_train'][:]
        print(f"Исходный train: {X_train.shape}")
        print(f"Распределение классов: {np.bincount(y_train)}")

        # Тестируем аугментацию
        augmenter = TimeSeriesAugmenter(
            noise_std=0.005,
            warping_sigma=0.2,
            scaling_range=(0.95, 1.05),
            mixup_alpha=0.2,
            dropout_p=0.1,
        )

        X_aug, y_aug = augmenter.augment(
            X_train, y_train,
            methods=['noise', 'warp', 'scale', 'dropout', 'mixup'],
        )
        print(f"Аугментированный train: {X_aug.shape}")
        print(f"  - Noise: +{len(X_train)}")
        print(f"  - Warp:  +{len(X_train)}")
        print(f"  - Scale: +{len(X_train)}")
        print(f"  - Dropout: +{len(X_train)}")
        print(f"  - Mixup: +{len(X_train)}")
        print(f"  Итого: {len(X_aug)} = 6× исходных данных")
    else:
        print(f"Файл {h5_path} не найден. Сначала создайте датасет через pipeline.py")
