"""
MultiTimeframeFusion Dataset — Alignment H1 + D1 по timestamp.

H1 и D1 имеют разную частоту (H1 ~8 баров/день, D1 ~1 бар/день).
Алгоритм: для каждого D1-бара (точка принятия решения) собираем:
  - D1_seq: последние d1_seq_len D1-баров (дневной контекст)
  - H1_seq: последние h1_seq_len H1-баров, timestamp <= D1_timestamp

Цель: направление следующего D1-бара (BUY/HOLD/SELL).
"""

from typing import Optional, Tuple

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

from src.db.connection import fetch_ohlcv_combined
from src.ml.features.price_features import add_price_features
from src.ml.features.indicator_features import add_indicator_features
from src.ml.features.calendar_features import add_calendar_features
from src.ml.features.target_encoding import encode_target


class MultiTimeframeFusionDataset(Dataset):
    """
    Dataset для MultiTimeframeFusion: H1 (intraday) + D1 (daily).

    Каждый сэмпл = (h1_seq, d1_seq, target), где:
      - h1_seq: (h1_seq_len, n_features) — последние H1-бары
      - d1_seq: (d1_seq_len, n_features) — последние D1-бары
      - target: 0=HOLD, 1=BUY, 2=SELL

    Args:
        ticker: тикер инструмента.
        h1_limit: сколько H1-баров загружать (≥ h1_seq_len + запас).
        d1_limit: сколько D1-баров загружать (≥ d1_seq_len + запас).
        h1_seq_len: длина окна H1.
        d1_seq_len: длина окна D1.
        h1_window_days: сколько торговых дней H1-данных включать (для грубой оценки).
        forecast_horizon: горизонт прогноза в D1-барах.
        threshold_pct: порог доходности для классов BUY/SELL (%%).
        val_split: доля валидации (хронологически, от未来 к прошлому).
        test_split: доля теста.
        use_tbank: использовать T-Bank API для свежих данных.
    """

    TARGET_NAMES = {0: 'HOLD', 1: 'BUY', 2: 'SELL'}

    def __init__(
        self,
        ticker: str = 'X5',
        h1_limit: int = 3000,
        d1_limit: int = 2000,
        h1_seq_len: int = 120,
        d1_seq_len: int = 30,
        forecast_horizon: int = 1,
        threshold_pct: float = 0.5,
        val_split: float = 0.15,
        test_split: float = 0.15,
        use_tbank: bool = True,
        feature_cols: Optional[list] = None,
        top_k_features: Optional[int] = None,
    ):
        self.ticker = ticker
        self.h1_seq_len = h1_seq_len
        self.d1_seq_len = d1_seq_len
        self.forecast_horizon = forecast_horizon
        self.threshold_pct = threshold_pct
        self.val_split = val_split
        self.test_split = test_split
        self.top_k_features = top_k_features

        # 1. Загрузка H1 и D1 данных
        print(f"[MTF Dataset] Загрузка {ticker} H1 (лимит={h1_limit})...")
        h1_df = fetch_ohlcv_combined(ticker, 'H1', limit=h1_limit, use_tbank=use_tbank)
        print(f"  H1: {len(h1_df)} баров, {h1_df['Date'].iloc[0]} — {h1_df['Date'].iloc[-1]}")

        print(f"[MTF Dataset] Загрузка {ticker} D1 (лимит={d1_limit})...")
        d1_df = fetch_ohlcv_combined(ticker, 'D1', limit=d1_limit, use_tbank=use_tbank)
        print(f"  D1: {len(d1_df)} баров, {d1_df['Date'].iloc[0]} — {d1_df['Date'].iloc[-1]}")

        # 2. Фичи для H1
        print("[MTF Dataset] Feature engineering H1...")
        h1_df = add_price_features(h1_df.copy())
        h1_df = add_indicator_features(h1_df)
        h1_df = add_calendar_features(h1_df)

        # 3. Фичи для D1
        print("[MTF Dataset] Feature engineering D1...")
        d1_df = add_price_features(d1_df.copy())
        d1_df = add_indicator_features(d1_df)
        d1_df = add_calendar_features(d1_df)

        # 4. Цель на D1
        print("[MTF Dataset] Encoding target...")
        d1_df = encode_target(
            d1_df,
            forecast_horizon=forecast_horizon,
            method='direction',
            threshold_pct=threshold_pct,
        )

        # 5. Forward-fill + удаление NaN для обоих датафреймов
        h1_df = h1_df.ffill().bfill().dropna()
        d1_df = d1_df.ffill().bfill().dropna(subset=['target'])
        # Гарантируем отсутствие inf
        h1_df = h1_df.replace([np.inf, -np.inf], 0.0)
        d1_df = d1_df.replace([np.inf, -np.inf], 0.0)

        # Определяем колонки признаков (все, кроме meta и target)
        if feature_cols is not None:
            # Используем переданный список (для cross-ticker consistency)
            self.feature_cols = [c for c in feature_cols if c in d1_df.columns]
            # Проверяем, что все колонки нашлись
            missing = set(feature_cols) - set(d1_df.columns)
            if missing:
                print(f"  ⚠️ Отсутствуют колонки: {missing}")
        else:
            exclude_cols = {'timestamp', 'Date', 'Time', 'target'}
            self.feature_cols = [
                c for c in d1_df.columns
                if c not in exclude_cols and not c.startswith('target_')
                and d1_df[c].nunique() > 1  # убираем константные
            ]
        print(f"  Признаков до FS: {len(self.feature_cols)}")

        # 6. Feature Selection (опционально)
        if self.top_k_features is not None and self.top_k_features < len(self.feature_cols):
            print(f"[MTF Dataset] Feature selection: {len(self.feature_cols)} → {self.top_k_features}")
            from sklearn.feature_selection import mutual_info_classif
            X_fs = d1_df[self.feature_cols].values.astype(np.float64)
            X_fs = np.nan_to_num(X_fs, nan=0.0, posinf=10.0, neginf=-10.0)
            y_fs = d1_df['target'].values
            mi_scores = mutual_info_classif(X_fs, y_fs, random_state=42)
            sorted_idx = np.argsort(mi_scores)[::-1]
            top_idx = sorted_idx[:self.top_k_features]
            # Сохраняем feature_cols но с фильтром
            self.feature_cols = [self.feature_cols[i] for i in top_idx]
            print(f"  Топ-5: {self.feature_cols[:5]}")
            # Обновляем массивы признаков (уже загружены ниже, но колонки обновлены)
            print(f"  Признаков после FS: {len(self.feature_cols)}")

        # 7. Alignment H1 → D1 по timestamp
        # Для каждого D1-бара находим H1-бары с timestamp <= D1_timestamp
        # и берём последние h1_seq_len из них
        print("[MTF Dataset] Aligning H1 → D1 by timestamps...")

        h1_timestamps = h1_df['timestamp'].values
        d1_timestamps = d1_df['timestamp'].values
        h1_features = h1_df[self.feature_cols].values.astype(np.float64)
        d1_features = d1_df[self.feature_cols].values.astype(np.float64)
        d1_targets = d1_df['target'].values.astype(np.int64)

        # Проверка на NaN/Inf в features
        for name, arr in [('H1', h1_features), ('D1', d1_features)]:
            nan_count = int(np.isnan(arr).sum())
            inf_count = int(np.isinf(arr).sum())
            if nan_count > 0 or inf_count > 0:
                print(f"  ⚠️ {name}: NaN={nan_count}, Inf={inf_count} — заменяем на 0")
                arr = np.nan_to_num(arr, nan=0.0, posinf=10.0, neginf=-10.0)

        # Массивы для хранения выровненных сэмплов
        h1_seqs, d1_seqs, targets, valid_indices = [], [], [], []

        min_h1_needed = h1_seq_len + 10  # запас

        for i in range(len(d1_df)):
            d1_ts = d1_timestamps[i]

            # D1 окно
            d1_start = i - d1_seq_len + 1
            if d1_start < 0:
                continue  # недостаточно D1 истории

            # H1 окно: все H1-бары с timestamp <= D1_timestamp
            h1_mask = h1_timestamps <= d1_ts
            h1_count = h1_mask.sum()
            if h1_count < min_h1_needed:
                continue  # недостаточно H1 истории

            # Берём последние h1_seq_len H1-баров
            h1_end = h1_count  # последний индекс, удовлетворяющий условию
            h1_start = max(0, h1_end - h1_seq_len)
            actual_h1_len = h1_end - h1_start

            # Таргет: убеждаемся, что target не NaN
            target = d1_targets[i]
            if np.isnan(target) or target < 0 or target > 2:
                continue

            # Сохраняем
            h1_seqs.append(h1_features[h1_start:h1_end])
            d1_seqs.append(d1_features[d1_start:i+1])
            targets.append(target)
            valid_indices.append(i)

        if not h1_seqs:
            raise ValueError("Нет ни одного валидного сэмпла после alignment!")

        # 7. Преобразуем в numpy (padding H1 если длина меньше h1_seq_len)
        X_h1 = np.zeros((len(h1_seqs), h1_seq_len, len(self.feature_cols)), dtype=np.float64)
        X_d1 = np.stack(d1_seqs)
        y = np.array(targets, dtype=np.int64)

        for idx, h1_seq in enumerate(h1_seqs):
            actual_len = len(h1_seq)
            if actual_len < h1_seq_len:
                X_h1[idx, h1_seq_len - actual_len:] = h1_seq
            else:
                X_h1[idx] = h1_seq[-h1_seq_len:]

        self.valid_indices = valid_indices
        self.d1_dates = d1_df['Date'].values[valid_indices]

        # 8. Хронологический split
        n = len(X_h1)
        test_n = int(n * test_split)
        val_n = int(n * val_split)

        train_end = n - test_n - val_n
        val_end = n - test_n

        self.data = {
            'train': {
                'h1': X_h1[:train_end],
                'd1': X_d1[:train_end],
                'y': y[:train_end],
                'dates': self.d1_dates[:train_end],
            },
            'val': {
                'h1': X_h1[train_end:val_end],
                'd1': X_d1[train_end:val_end],
                'y': y[train_end:val_end],
                'dates': self.d1_dates[train_end:val_end],
            },
            'test': {
                'h1': X_h1[val_end:],
                'd1': X_d1[val_end:],
                'y': y[val_end:],
                'dates': self.d1_dates[val_end:],
            },
        }

        # Печать статистики
        self._print_stats()

    def _print_stats(self):
        """Вывести статистику датасета."""
        print(f"\n[MTF Dataset] СТАТИСТИКА:")
        print(f"  Всего сэмплов: {len(self.data['train']['y']) + len(self.data['val']['y']) + len(self.data['test']['y'])}")
        for split_name in ['train', 'val', 'test']:
            d = self.data[split_name]
            dist = {self.TARGET_NAMES[k]: v for k, v in
                    sorted(zip(*np.unique(d['y'], return_counts=True)))}
            print(f"  {split_name.capitalize():5s}: {len(d['y']):5d} сэмплов, "
                  f"H1={d['h1'].shape}, D1={d['d1'].shape}, "
                  f"распред: {dist}")

    def get_split(self, split: str = 'train') -> 'MultiTimeframeFusionSubset':
        """Получить подмножество для train/val/test."""
        if split not in self.data:
            raise ValueError(f"Split '{split}' не найден. Доступны: train, val, test")
        return MultiTimeframeFusionSubset(self.data[split])

    def get_feature_cols(self) -> list:
        """Вернуть список названий признаков."""
        return self.feature_cols


class MultiTimeframeFusionSubset(Dataset):
    """
    Подмножество MTF датасета для использования с DataLoader.

    Каждый сэмпл: (h1_tensor, d1_tensor, target)
    """

    def __init__(self, data: dict):
        self.h1 = data['h1']
        self.d1 = data['d1']
        self.y = data['y']
        self.dates = data.get('dates', [])

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

    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        return (
            torch.FloatTensor(self.h1[idx]),
            torch.FloatTensor(self.d1[idx]),
            torch.LongTensor([self.y[idx]])[0],
        )
