"""
Пайплайн подготовки ML-датасета из MOEX данных.

Функции:
    create_ml_dataset — полный цикл: загрузка → фичи → таргет → сплит → нормализация → последовательности → сохранение
    create_dataset_for_training — упрощённый интерфейс для DataLoader
    save_dataset_to_hdf5 — сохранение датасета в HDF5
    print_dataset_metadata — вывод метаданных датасета
"""

import json
import os
import sys
from pathlib import Path
from typing import Any, Optional, Union

import h5py
import numpy as np
import pandas as pd

# Добавляем корень проекта в PYTHONPATH
PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))

from sklearn.preprocessing import StandardScaler

from src.db.connection import fetch_ohlcv_combined
from src.ml.data.dataset import TimeSeriesDataset, create_sequences
from src.ml.data.quality import run_data_quality_checks, validate_dataset
from src.ml.features.calendar_features import add_calendar_features
from src.ml.features.indicator_features import add_indicator_features
from src.ml.features.price_features import add_price_features
from src.ml.features.target_encoding import encode_target


def temporal_train_val_test_split(
    df: pd.DataFrame,
    val_size: float = 0.15,
    test_size: float = 0.15,
) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:
    """
    Хронологическое разбиение на train/val/test (без перемешивания).

    Важно: временные ряды НЕЛЬЗЯ перемешивать случайно.
    Разбиение идёт от прошлого к будущему: train → val → test.

    Args:
        df: DataFrame, отсортированный по возрастанию timestamp.
        val_size: доля валидационной выборки.
        test_size: доля тестовой выборки.

    Returns:
        Кортеж (train_df, val_df, test_df).
    """
    n = len(df)
    test_start = int(n * (1 - test_size))
    val_start = int(n * (1 - test_size - val_size))

    train = df.iloc[:val_start].copy()
    val = df.iloc[val_start:test_start].copy()
    test = df.iloc[test_start:].copy()

    print(f"  Разбиение: train={len(train)} ({len(train)/n*100:.1f}%), "
          f"val={len(val)} ({len(val)/n*100:.1f}%), "
          f"test={len(test)} ({len(test)/n*100:.1f}%)")

    return train, val, test


def create_ml_dataset(
    ticker: str,
    tf: str = 'D1',
    seq_len: int = 60,
    forecast_horizon: int = 1,
    threshold_pct: float = 0.5,
    target_method: str = 'direction',
    val_split: float = 0.15,
    test_split: float = 0.15,
    limit: int = 2000,
    save_path: Optional[str] = None,
) -> dict[str, Any]:
    """
    Полный пайплайн создания ML-датасета.

    Этапы:
        1. Загрузка данных из БД (MySQL + T-Bank API)
        2. Добавление price-признаков (доходности, OHLC-отношения)
        3. Добавление индикаторов (EMA, MACD, RSI, Stochastic, ADX, BB, ATR, OBV, MFI)
        4. Добавление календарных признаков (день недели, месяц, sin/cos)
        5. Кодирование целевой переменной (3 класса: BUY/HOLD/SELL)
        6. Удаление строк с NaN
        7. Хронологическое разбиение train/val/test
        8. Нормализация (Z-score, fit на train)
        9. Создание скользящих окон (seq_len)
        10. Сохранение в HDF5

    Args:
        ticker: тикер инструмента (например, 'X5').
        tf: таймфрейм ('H1', 'D1', 'W1').
        seq_len: длина окна (количество прошлых наблюдений).
        forecast_horizon: горизонт прогноза (свечей вперёд).
        threshold_pct: порог доходности в процентах для BUY/SELL.
        target_method: метод кодирования цели ('direction', 'binary', 'regression').
        val_split: доля валидационной выборки.
        test_split: доля тестовой выборки.
        limit: максимальное количество загружаемых свечей.
        save_path: путь для сохранения HDF5 файла (если None — не сохранять).

    Returns:
        Словарь с результатами:
            - 'train_dataset': TimeSeriesDataset
            - 'val_dataset': TimeSeriesDataset
            - 'test_dataset': TimeSeriesDataset
            - 'scaler': StandardScaler (обучен на train)
            - 'feature_names': список названий признаков
            - 'metadata': словарь с метаданными
    """
    print(f"=" * 70)
    print(f"  ПАЙПЛАЙН ПОДГОТОВКИ ML-ДАТАСЕТА")
    print(f"  Тикер: {ticker} | ТФ: {tf} | seq_len: {seq_len} | горизонт: {forecast_horizon}")
    print(f"=" * 70)

    # ── 1. Загрузка данных ─────────────────────────────────────────
    print(f"\n[1/10] Загрузка данных {ticker}_{tf} из БД...")
    df = fetch_ohlcv_combined(ticker, tf, limit=limit)
    print(f"  Загружено {len(df)} свечей, "
          f"диапазон: {df['Date'].iloc[0]} — {df['Date'].iloc[-1]}")

    # ── 2. Price-признаки ──────────────────────────────────────────
    print(f"\n[2/10] Добавление price-признаков...")
    df = add_price_features(df, return_horizons=[1, 5, 10, 21])
    print(f"  Добавлены: доходности (1,5,10,21), log_ret, OHLC-отношения, close_position")
    print(f"  Размер: {df.shape}")

    # ── 3. Индикаторы ─────────────────────────────────────────────
    print(f"\n[3/10] Добавление технических индикаторов...")
    df = add_indicator_features(
        df,
        include_trend=True,
        include_oscillators=True,
        include_volatility=True,
        include_volume=True,
    )
    print(f"  Добавлены: EMA(9,21,50), SMA(20,50,200), MACD, ADX, "
          f"RSI(7,14), Stoch, CCI, MFI, BB, ATR, OBV, VolumeSMA")
    print(f"  Размер: {df.shape}")

    # ── 4. Календарные признаки ───────────────────────────────────
    print(f"\n[4/10] Добавление календарных признаков...")
    df = add_calendar_features(
        df,
        include_datetime=True,
        include_cyclical=True,
        include_sessions=True,
    )
    print(f"  Добавлены: день_недели, месяц, квартал, sin/cos, сессии")
    print(f"  Размер: {df.shape}")

    # ── 5. Целевая переменная ──────────────────────────────────────
    print(f"\n[5/10] Кодирование целевой переменной (метод: {target_method}, "
          f"порог: {threshold_pct}%)...")
    df = encode_target(
        df,
        forecast_horizon=forecast_horizon,
        method=target_method,
        threshold_pct=threshold_pct,
    )
    target_dist = df['target'].value_counts()
    print(f"  Распределение цели:\n    {target_dist.to_dict()}")

    # ── 6. Качество данных ─────────────────────────────────────────
    print(f"\n[6/10] Проверка качества данных...")
    quality_report = run_data_quality_checks(df, target_col='target')
    print(f"  Строк: {quality_report['shape']['rows']}, колонок: {quality_report['shape']['cols']}")
    print(f"  Пропуски: {'✅ OK' if not quality_report['missing_values']['has_nulls'] else '⚠️ есть'}")
    print(f"  Выбросы: {'✅ OK' if not quality_report['outliers']['has_outliers'] else '⚠️ есть'}")
    print(f"  Баланс классов: {'✅ OK' if not quality_report['class_balance']['is_imbalanced'] else '⚠️ дисбаланс'}")
    print(f"  Хронология: {'✅ OK' if quality_report['temporal_order'].get('is_sorted') else '⚠️ нарушена'}")
    print(f"  Константные фичи: {quality_report['constant_features']['constant_features']}")

    # ── Обработка NaN ──────────────────────────────────────────
    print(f"\n  Обработка NaN: forward-fill + удаление оставшихся...")
    # Запоминаем исходное количество
    before_drop = len(df)
    # Forward-fill для индикаторов (первые N строк будут NaN)
    df = df.ffill()
    # Удаляем оставшиеся NaN
    df = df.dropna()
    # Удаляем строки с NaN в target
    df = df.dropna(subset=['target'])
    after_drop = len(df)
    print(f"  Удалено {before_drop - after_drop} строк с NaN")
    print(f"  Итоговый размер: {df.shape}")

    if after_drop < seq_len + 100:
        raise ValueError(
            f"Слишком мало данных после очистки: {after_drop}. "
            f"Нужно хотя бы {seq_len + 100} для осмысленного датасета"
        )

    # ── 7. Определение колонок признаков ─────────────────────────
    exclude_cols = {
        'timestamp', 'Date', 'Time', 'target',
    }
    feature_cols = [
        c for c in df.columns
        if c not in exclude_cols and not c.startswith('target_')
    ]
    # Автоматически удаляем константные признаки (если есть)
    constant_feats = set()
    for c in feature_cols:
        if df[c].nunique() <= 1:
            constant_feats.add(c)
    if constant_feats:
        print(f"  ⚠ Удаляем константные признаки: {constant_feats}")
        feature_cols = [c for c in feature_cols if c not in constant_feats]
    print(f"\n[7/10] Выбрано {len(feature_cols)} признаков:")
    for i, f in enumerate(feature_cols[:10]):
        print(f"    {i+1}. {f}")
    if len(feature_cols) > 10:
        print(f"    ... и ещё {len(feature_cols) - 10}")

    # ── 8-9. Нормализация (Z-score, fit на train) + Split ──────────
    # Новый подход: сначала нормализуем, потом создаём последовательности,
    # потом хронологически split-им последовательности.
    # Это гарантирует, что все сплиты имеют достаточно сэмплов.
    print(f"\n[8-9] Нормализация + хронологическое разбиение...")

    # Разбиваем сырые данные для нормализации
    n = len(df)
    test_start_raw = int(n * (1 - test_split))
    val_start_raw = int(n * (1 - test_split - val_split))

    # Normalizer fit только на train части
    scaler = StandardScaler()
    train_raw_data = df.iloc[:val_start_raw][feature_cols].values
    scaler.fit(train_raw_data)
    print(f"  Normalizer fit на train ({len(train_raw_data)} строк)")

    # Трансформируем все данные
    X_all = scaler.transform(df[feature_cols].values)
    y_all = df['target'].values.astype(np.int64)

    print(f"  Средние (первые 5): {scaler.mean_[:5].round(3)}")
    print(f"  Std (первые 5): {scaler.scale_[:5].round(3)}")

    # ── 10. Создание последовательностей + split ──────────────────
    print(f"\n[10/10] Создание скользящих окон (seq_len={seq_len}) и хронологическое разбиение...")
    X_seq, y_seq = create_sequences(X_all, y_all, seq_len)

    # Хронологически разбиваем последовательности
    n_seq = len(X_seq)
    test_seq_start = int(n_seq * (1 - test_split))
    val_seq_start = int(n_seq * (1 - test_split - val_split))

    X_train_seq = X_seq[:val_seq_start]
    y_train_seq = y_seq[:val_seq_start]
    X_val_seq = X_seq[val_seq_start:test_seq_start]
    y_val_seq = y_seq[val_seq_start:test_seq_start]
    X_test_seq = X_seq[test_seq_start:]
    y_test_seq = y_seq[test_seq_start:]

    print(f"  Train: X {X_train_seq.shape}, y {y_train_seq.shape}")
    print(f"  Val:   X {X_val_seq.shape}, y {y_val_seq.shape}")
    print(f"  Test:  X {X_test_seq.shape}, y {y_test_seq.shape}")

    # Даты (последний элемент окна = seq_len-1 + idx)
    seq_offset = seq_len - 1
    date_series = df['Date'].values
    if len(X_train_seq) > 0:
        print(f"  Train даты: {date_series[seq_offset]} — {date_series[seq_offset + len(X_train_seq) - 1]}")
    if len(X_val_seq) > 0:
        vi = seq_offset + val_seq_start
        print(f"  Val даты:   {date_series[vi]} — {date_series[vi + len(X_val_seq) - 1]}")
    if len(X_test_seq) > 0:
        ti = seq_offset + test_seq_start
        print(f"  Test даты:  {date_series[ti]} — {date_series[ti + len(X_test_seq) - 1]}")

    # Проверка на NaN
    for name, arr in [('train', X_train_seq), ('val', X_val_seq), ('test', X_test_seq)]:
        n_nan = int(np.isnan(arr).sum())
        if n_nan > 0:
            print(f"  ⚠️ NaN в {name}: {n_nan}")
            np.nan_to_num(arr, copy=False, nan=0.0)
    print(f"  ✅ NaN: 0")

    # Создание PyTorch Dataset (seq_len=1, т.к. окна уже созданы)
    train_dataset = TimeSeriesDataset(X_train_seq, y_train_seq, seq_len=1) if len(X_train_seq) > 0 else None
    val_dataset = TimeSeriesDataset(X_val_seq, y_val_seq, seq_len=1) if len(X_val_seq) > 0 else None
    test_dataset = TimeSeriesDataset(X_test_seq, y_test_seq, seq_len=1) if len(X_test_seq) > 0 else None

    # Валидация
    metadata = validate_dataset(
        X_train_seq, y_train_seq,
        X_val_seq, y_val_seq,
        X_test_seq, y_test_seq,
        feature_names=feature_cols,
    )
    metadata['ticker'] = ticker
    metadata['tf'] = tf
    metadata['seq_len'] = seq_len
    metadata['forecast_horizon'] = forecast_horizon
    metadata['threshold_pct'] = threshold_pct
    metadata['target_method'] = target_method
    metadata['date_range'] = {
        'start': str(df['Date'].iloc[0]),
        'end': str(df['Date'].iloc[-1]),
    }
    metadata['total_raw_rows'] = int(len(df))

    # ── Сохранение в HDF5 ───────────────────────────────────────────
    if save_path:
        save_dataset_to_hdf5(
            save_path=save_path,
            X_train=X_train_seq, y_train=y_train_seq,
            X_val=X_val_seq, y_val=y_val_seq,
            X_test=X_test_seq, y_test=y_test_seq,
            feature_names=feature_cols,
            scaler_mean=scaler.mean_,
            scaler_scale=scaler.scale_,
            metadata=metadata,
        )

    return {
        'train_dataset': train_dataset,
        'val_dataset': val_dataset,
        'test_dataset': test_dataset,
        'X_train': X_train_seq,
        'y_train': y_train_seq,
        'X_val': X_val_seq,
        'y_val': y_val_seq,
        'X_test': X_test_seq,
        'y_test': y_test_seq,
        'scaler': scaler,
        'feature_names': feature_cols,
        'metadata': metadata,
    }


def save_dataset_to_hdf5(
    save_path: str,
    X_train: np.ndarray,
    y_train: np.ndarray,
    X_val: np.ndarray,
    y_val: np.ndarray,
    X_test: np.ndarray,
    y_test: np.ndarray,
    feature_names: list[str],
    scaler_mean: np.ndarray,
    scaler_scale: np.ndarray,
    metadata: dict[str, Any],
) -> str:
    """
    Сохранить датасет в HDF5 файл.

    Структура HDF5:
        /X_train, /y_train, /X_val, /y_val, /X_test, /y_test — массивы
        /feature_names — названия признаков
        /scaler_mean, /scaler_scale — параметры нормализации
        /metadata — JSON с метаданными

    Args:
        save_path: путь для сохранения.
        X_train, y_train, X_val, y_val, X_test, y_test: массивы.
        feature_names: названия признаков.
        scaler_mean, scaler_scale: параметры StandardScaler.
        metadata: словарь метаданных.

    Returns:
        Путь к сохранённому файлу.
    """
    os.makedirs(os.path.dirname(save_path), exist_ok=True)

    with h5py.File(save_path, 'w') as f:
        f.create_dataset('X_train', data=X_train, compression='gzip', compression_opts=4)
        f.create_dataset('y_train', data=y_train, compression='gzip', compression_opts=4)
        f.create_dataset('X_val', data=X_val, compression='gzip', compression_opts=4)
        f.create_dataset('y_val', data=y_val, compression='gzip', compression_opts=4)
        f.create_dataset('X_test', data=X_test, compression='gzip', compression_opts=4)
        f.create_dataset('y_test', data=y_test, compression='gzip', compression_opts=4)

        # Feature names
        dt = h5py.string_dtype()
        f.create_dataset('feature_names', data=feature_names, dtype=dt)

        # Scaler params
        f.create_dataset('scaler_mean', data=scaler_mean)
        f.create_dataset('scaler_scale', data=scaler_scale)

        # Metadata as JSON
        metadata_json = json.dumps(metadata, indent=2, default=str)
        f.create_dataset('metadata', data=metadata_json, dtype=dt)

    file_size = os.path.getsize(save_path)
    print(f"\n  ✅ Датасет сохранён: {save_path}")
    print(f"  Размер файла: {file_size / 1024 / 1024:.2f} MB")

    return save_path


def print_dataset_metadata(metadata: dict[str, Any]) -> None:
    """
    Вывести метаданные датасета в читаемом формате.

    Args:
        metadata: словарь с метаданными из create_ml_dataset.
    """
    print(f"\n{'=' * 70}")
    print(f"  📊 МЕТАДАННЫЕ ДАТАСЕТА")
    print(f"{'=' * 70}")
    print(f"  Тикер:                {metadata.get('ticker', '—')}")
    print(f"  Таймфрейм:            {metadata.get('tf', '—')}")
    print(f"  Длина окна (seq_len): {metadata.get('seq_len', '—')}")
    print(f"  Горизонт прогноза:    {metadata.get('forecast_horizon', '—')}")
    print(f"  Порог классификации:  {metadata.get('threshold_pct', '—')}%")
    print(f"  Метод таргета:        {metadata.get('target_method', '—')}")
    print(f"  Диапазон дат:         {metadata.get('date_range', {}).get('start', '—')} — "
          f"{metadata.get('date_range', {}).get('end', '—')}")
    print(f"  Всего строк (сырых):  {metadata.get('total_raw_rows', '—')}")
    print(f"  Количество признаков: {metadata.get('n_features', '—')}")

    # Feature names
    feature_names = metadata.get('feature_names', [])
    if feature_names:
        print(f"  Признаки ({len(feature_names)}):")
        # Группируем
        price_feats = [f for f in feature_names if f.startswith(('ret_', 'log_ret', 'high_low', 'close_open',
                                                                    'upper_shadow', 'lower_shadow', 'spread', 'close_position'))]
        ind_feats = [f for f in feature_names if f.startswith(('ema_', 'sma_', 'price_to', 'macd', 'adx',
                                                                 'plus_di', 'minus_di', 'rsi_', 'stoch_',
                                                                 'cci_', 'mfi_', 'bb_', 'atr_', 'obv', 'volume'))]
        cal_feats = [f for f in feature_names if f.startswith(('day_', 'month_', 'quarter', 'is_', 'day_sin',
                                                                 'day_cos', 'month_sin', 'month_cos'))]
        if price_feats:
            print(f"    📈 Price ({len(price_feats)}): {', '.join(price_feats)}")
        if ind_feats:
            print(f"    🔧 Indicators ({len(ind_feats)}): {', '.join(ind_feats[:10])}...")
        if cal_feats:
            print(f"    📅 Calendar ({len(cal_feats)}): {', '.join(cal_feats)}")

    print(f"  Размеры:")
    print(f"    Train: {metadata.get('train_samples', '—')} сэмплов")
    print(f"    Val:   {metadata.get('val_samples', '—')} сэмплов")
    print(f"    Test:  {metadata.get('test_samples', '—')} сэмплов")

    # Баланс классов
    for split in ['train', 'val', 'test']:
        class_pct = metadata.get(f'class_pct_{split}', {})
        if class_pct:
            cls_str = ', '.join([
                f"{['HOLD','BUY','SELL'][int(k)] if int(k) in [0,1,2] else f'Class{k}'}: {v}%"
                for k, v in sorted(class_pct.items())
            ])
            print(f"    {split.capitalize()} классы: {cls_str}")

    print(f"  NaN после нормализации: {metadata.get('nan_count', '—')}")
    print(f"{'=' * 70}\n")


def create_dataset_for_training(
    ticker: str,
    tf: str = 'D1',
    seq_len: int = 60,
    forecast_horizon: int = 1,
    threshold_pct: float = 0.5,
    val_split: float = 0.15,
    test_split: float = 0.15,
    limit: int = 2000,
    save_dir: Optional[str] = None,
) -> dict[str, Any]:
    """
    Упрощённый интерфейс для создания датасета (совместимость).

    Args:
        ticker: тикер инструмента.
        tf: таймфрейм.
        seq_len: длина окна.
        forecast_horizon: горизонт прогноза.
        threshold_pct: порог доходности в процентах.
        val_split: доля валидации.
        test_split: доля теста.
        limit: максимальное количество свечей.
        save_dir: директория для сохранения (если None — не сохранять).

    Returns:
        Словарь с датасетами и метаданными.
    """
    save_path = None
    if save_dir:
        os.makedirs(save_dir, exist_ok=True)
        save_path = os.path.join(save_dir, f"{ticker.lower()}_{tf.lower()}_dataset.h5")

    result = create_ml_dataset(
        ticker=ticker,
        tf=tf,
        seq_len=seq_len,
        forecast_horizon=forecast_horizon,
        threshold_pct=threshold_pct,
        target_method='direction',
        val_split=val_split,
        test_split=test_split,
        limit=limit,
        save_path=save_path,
    )

    return result


if __name__ == '__main__':
    """
    Скрипт для создания ML-датасета X5 D1.
    
    Использование:
        python src/ml/features/pipeline.py
    """
    print("Создание ML-датасета для X5 D1...")
    
    result = create_ml_dataset(
        ticker='X5',
        tf='D1',
        seq_len=60,
        forecast_horizon=1,
        threshold_pct=0.5,
        target_method='direction',
        val_split=0.15,
        test_split=0.15,
        limit=2000,
        save_path=str(PROJECT_ROOT / 'src' / 'ml' / 'models' / 'saved' / 'x5_d1_dataset.h5'),
    )

    metadata = result['metadata']
    print_dataset_metadata(metadata)
    
    print("✅ Датасет успешно создан!")
