"""
Multi-Timeframe обучение Wyckoff-модели (H1 + D1 + W1).

Пайплайн:
  1. Загрузка H1/D1/W1 для всех тикеров
  2. Feature engineering + Wyckoff-разметка для каждого ТФ
  3. Alignment по D1-таймстемпам (для каждого D1-бара: контекст H1 и W1)
  4. Обучение WyckoffMultiTimeframe
  5. Сохранение + регистрация в ModelRegistry

Usage:
    python src/ml/train/train_wyckoff_mtf.py
"""

import argparse
import json
import logging
import os
import sys
import time
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple

import numpy as np
import pandas as pd
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, Dataset, TensorDataset
from torch.optim import AdamW

PROJECT_ROOT = Path(__file__).resolve().parent.parent.parent.parent
sys.path.insert(0, str(PROJECT_ROOT))

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.wyckoff_features import add_wyckoff_features
from src.ml.data.wyckoff_labeling import WyckoffLabeler
from src.ml.models.wyckoff import WyckoffMultiTimeframe, create_wyckoff_model
from src.ml.models.registry import ModelRegistry
from src.ml.train.trainer import (
    compute_class_weights,
    set_seed,
)
from src.ml.train.losses import FocalLoss
from sklearn.metrics import f1_score

logging.basicConfig(
    level=logging.INFO,
    format='%(asctime)s | %(name)s | %(levelname)s | %(message)s',
    datefmt='%H:%M:%S',
)
logger = logging.getLogger('train_wyckoff_mtf')

DEFAULT_TICKERS = [
    'SBER', 'GAZP', 'LKOH', 'MOEX', 'MTSS', 'NSVZ', 'NVTK',
    'PHOR', 'PLZL', 'ROSN', 'SBER', 'SNGSP', 'VTBR', 'X5',
]

# Исключаем тикеры без данных в БД
SKIP_TICKERS = {'MGNT', 'GMKN', 'ALRS'}

TICKERS = [t for t in DEFAULT_TICKERS if t not in SKIP_TICKERS]
TICKERS = list(dict.fromkeys(TICKERS))  # unique, preserve order


def load_and_label_ticker_mtf(
    ticker: str,
    tf: str,
    limit: int = 2000,
    use_simple_phases: bool = True,
) -> Optional[pd.DataFrame]:
    """
    Загрузить данные тикера для ТФ, добавить признаки и разметку фаз.
    """
    try:
        df = fetch_ohlcv_combined(ticker, tf, limit=limit)
        if df is None or df.empty:
            logger.warning(f"  {ticker}_{tf}: нет данных")
            return None

        # Убеждаемся, что колонки числовые
        for c in ['Open', 'High', 'Low', 'Close', 'Volume']:
            df[c] = pd.to_numeric(df[c], errors='coerce')
        df = df.dropna(subset=['Close'])

        if len(df) < 100:
            logger.warning(f"  {ticker}_{tf}: слишком мало данных ({len(df)})")
            return None

        # Feature engineering
        df = add_price_features(df, return_horizons=[1, 5, 10, 21])
        df = add_indicator_features(df, include_trend=True, include_oscillators=True,
                                     include_volatility=True, include_volume=True)
        df = add_wyckoff_features(df)

        # Wyckoff-разметка
        labeler = WyckoffLabeler(smooth_window=5)
        df = labeler.label_phases(df)

        target_col = 'wyckoff_phase_simple' if use_simple_phases else 'wyckoff_phase'
        df = df.ffill().dropna(subset=[target_col])

        phase_dist = df[target_col].value_counts().sort_index()
        logger.info(f"  {ticker}_{tf}: {len(df)} строк, фазы: {dict(phase_dist)}")

        return df

    except Exception as e:
        logger.error(f"  {ticker}_{tf}: ошибка: {e}")
        return None


def get_feature_cols(df: pd.DataFrame) -> List[str]:
    """Получить список колонок-признаков из датафрейма."""
    exclude_cols = {
        'timestamp', 'Date', 'Time', 'Open', 'High', 'Low', 'Close', 'Volume',
        'target', 'wyckoff_phase', 'wyckoff_phase_name',
        'wyckoff_phase_simple', 'wyckoff_phase_simple_name',
        'hh_hl', 'lh_ll',
    }
    feature_cols = [c for c in df.columns if c not in exclude_cols
                    and not c.startswith('target_')
                    and df[c].dtype in [np.float64, np.float32, np.int64, np.int32]]
    # Удаляем константные
    for c in feature_cols[:]:
        if df[c].nunique() <= 1:
            feature_cols.remove(c)
    return feature_cols


def load_all_tickers_mtf(
    tickers: List[str],
    limit: int = 2000,
    use_simple_phases: bool = True,
) -> Dict[str, Dict[str, Any]]:
    """
    Загрузить все тикеры для всех ТФ.
    
    Returns:
        {ticker: {'H1': df, 'D1': df, 'W1': df}, ...}
    """
    result = {}
    for ticker in tickers:
        ticker_data = {}
        for tf in ['H1', 'D1', 'W1']:
            logger.info(f"Загрузка {ticker}_{tf}...")
            df = load_and_label_ticker_mtf(ticker, tf, limit, use_simple_phases)
            if df is not None:
                ticker_data[tf] = df
        
        if 'D1' in ticker_data and 'H1' in ticker_data and 'W1' in ticker_data:
            result[ticker] = ticker_data
        else:
            missing = [tf for tf in ['H1', 'D1', 'W1'] if tf not in ticker_data]
            logger.warning(f"  {ticker}: пропущен из-за отсутствия {missing}")

    logger.info(f"\nЗагружено тикеров: {len(result)}")
    return result


def compute_global_features(
    all_data: Dict[str, Dict[str, pd.DataFrame]],
    tf_order: List[str] = None,
) -> List[str]:
    """
    Вычислить пересечение признаков по всем тикерам и ТФ.
    """
    if tf_order is None:
        tf_order = ['H1', 'D1', 'W1']
    
    global_cols = None
    for ticker, tfs in all_data.items():
        for tf in tf_order:
            df = tfs.get(tf)
            if df is None:
                continue
            cols = set(get_feature_cols(df))
            if global_cols is None:
                global_cols = cols
            else:
                global_cols = global_cols & cols
    
    if global_cols is None:
        raise ValueError("Нет общих признаков!")
    
    result = sorted(global_cols)
    logger.info(f"Глобальных признаков (пересечение всех ТФ+тикеров): {len(result)}")
    return result


def build_mtf_dataset(
    all_data: Dict[str, Dict[str, pd.DataFrame]],
    feature_cols: List[str],
    tf_config: Dict[str, int] = None,
    use_simple_phases: bool = True,
    val_split: float = 0.15,
    test_split: float = 0.15,
) -> Dict[str, Any]:
    """
    Создать MTF датасет: для каждого D1-бара алайним H1 и W1 контекст.
    
    Алаймент: для D1-бара с timestamp T берём:
      - H1: последние h1_len H1-баров с timestamp <= T
      - D1: последние d1_len D1-баров (включая T)
      - W1: последние w1_len W1-баров с timestamp <= T
      - target: фаза Вайкоффа на D1-баре T
    
    Возвращает словарь с данными для обучения.
    """
    if tf_config is None:
        tf_config = {'H1': 60, 'D1': 30, 'W1': 10}
    
    target_col = 'wyckoff_phase_simple' if use_simple_phases else 'wyckoff_phase'
    num_classes = 5 if use_simple_phases else 8
    n_features = len(feature_cols)
    
    all_h1, all_d1, all_w1, all_y = [], [], [], []
    all_tickers, all_dates = [], []
    
    for ticker, tfs in all_data.items():
        h1_df = tfs['H1']
        d1_df = tfs['D1']
        w1_df = tfs['W1']
        
        h1_ts = h1_df['timestamp'].values
        d1_ts = d1_df['timestamp'].values
        w1_ts = w1_df['timestamp'].values
        
        h1_feat = h1_df[feature_cols].values.astype(np.float32)
        d1_feat = d1_df[feature_cols].values.astype(np.float32)
        w1_feat = w1_df[feature_cols].values.astype(np.float32)
        d1_y = d1_df[target_col].values.astype(np.int64)
        
        h1_len, d1_len, w1_len = tf_config['H1'], tf_config['D1'], tf_config['W1']
        min_h1_needed = h1_len + 10
        
        ticker_h1, ticker_d1, ticker_w1, ticker_y = [], [], [], []
        ticker_dates = []
        
        for i in range(len(d1_df)):
            ts = d1_ts[i]
            
            # D1 окно
            d1_start = i - d1_len + 1
            if d1_start < 0:
                continue
            
            # H1 окно
            h1_mask = h1_ts <= ts
            h1_count = h1_mask.sum()
            if h1_count < min_h1_needed:
                continue
            
            # W1 окно
            w1_mask = w1_ts <= ts
            w1_count = w1_mask.sum()
            if w1_count < w1_len:
                continue
            
            # Берём последние N баров
            h1_end = h1_count
            h1_start = h1_end - h1_len
            
            w1_end = w1_count
            w1_start = w1_end - w1_len
            
            target = d1_y[i]
            if target < 0:
                continue
            
            ticker_h1.append(h1_feat[h1_start:h1_end])
            ticker_d1.append(d1_feat[d1_start:i+1])
            ticker_w1.append(w1_feat[w1_start:w1_end])
            ticker_y.append(target)
            ticker_dates.append(d1_df['Date'].iloc[i])
        
        if ticker_y:
            X_h1 = np.stack(ticker_h1)
            X_d1 = np.stack(ticker_d1)
            X_w1 = np.stack(ticker_w1)
            y = np.array(ticker_y, dtype=np.int64)
            
            all_h1.append(X_h1)
            all_d1.append(X_d1)
            all_w1.append(X_w1)
            all_y.append(y)
            all_tickers.extend([ticker] * len(y))
            all_dates.extend(ticker_dates)
            
            logger.info(f"  {ticker}: {len(ticker_y)} сэмплов")
    
    if not all_h1:
        raise ValueError("Нет данных для обучения!")
    
    # Конкатенация всех тикеров
    X_h1 = np.concatenate(all_h1, axis=0)
    X_d1 = np.concatenate(all_d1, axis=0)
    X_w1 = np.concatenate(all_w1, axis=0)
    y = np.concatenate(all_y, axis=0)
    
    logger.info(f"\nMTF Dataset: {len(y)} сэмплов, {n_features} признаков")
    logger.info(f"  H1: {X_h1.shape}, D1: {X_d1.shape}, W1: {X_w1.shape}")
    
    # Хронологический split
    n = len(y)
    test_n = int(n * test_split)
    val_n = int(n * val_split)
    
    train_end = n - test_n - val_n
    val_end = n - test_n
    
    data = {
        'X_h1_train': X_h1[:train_end],
        'X_d1_train': X_d1[:train_end],
        'X_w1_train': X_w1[:train_end],
        'y_train': y[:train_end],
        'X_h1_val': X_h1[train_end:val_end],
        'X_d1_val': X_d1[train_end:val_end],
        'X_w1_val': X_w1[train_end:val_end],
        'y_val': y[train_end:val_end],
        'X_h1_test': X_h1[val_end:],
        'X_d1_test': X_d1[val_end:],
        'X_w1_test': X_w1[val_end:],
        'y_test': y[val_end:],
        'feature_cols': feature_cols,
        'num_classes': num_classes,
        'n_features': n_features,
        'tf_config': tf_config,
    }
    
    # Статистика
    for split_name, Xh, Xd, Xw, yy in [
        ('Train', data['X_h1_train'], data['X_d1_train'], data['X_w1_train'], data['y_train']),
        ('Val', data['X_h1_val'], data['X_d1_val'], data['X_w1_val'], data['y_val']),
        ('Test', data['X_h1_test'], data['X_d1_test'], data['X_w1_test'], data['y_test']),
    ]:
        dist = {int(k): int(v) for k, v in zip(*np.unique(yy, return_counts=True))}
        logger.info(f"  {split_name}: {len(yy)} сэмплов, H1={Xh.shape}, распределение: {dist}")
    
    return data


def train_mtf_model(
    data: Dict[str, Any],
    n_epochs: int = 100,
    learning_rate: float = 1e-3,
    weight_decay: float = 1e-3,
    batch_size: int = 32,
    patience: int = 30,
    dropout: float = 0.4,
    model_name: str = 'wyckoff_mt_d1_mtf_v1',
    save_dir: str = 'src/ml/models/saved',
    use_focal_loss: bool = True,
    focal_gamma: float = 2.0,
) -> Tuple[nn.Module, Dict[str, Any]]:
    """
    Обучить WyckoffMultiTimeframe модель.
    """
    n_features = data['n_features']
    num_classes = data['num_classes']
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    
    logger.info(f"{'=' * 70}")
    logger.info(f"ОБУЧЕНИЕ WYCKOFF MTF")
    logger.info(f"  features={n_features}, num_classes={num_classes}")
    logger.info(f"  dropout={dropout}, batch_size={batch_size}")
    logger.info(f"  lr={learning_rate}, weight_decay={weight_decay}")
    logger.info(f"  patience={patience}, focal_gamma={focal_gamma}")
    logger.info(f"  device={device}")
    logger.info(f"{'=' * 70}")
    
    # Создание модели
    model = WyckoffMultiTimeframe(
        input_size=n_features,
        d_model=64,
        nhead=4,
        num_classes=num_classes,
        dropout=dropout,
    )
    model = model.to(device)
    
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    logger.info(f"  Параметры: всего {total_params:,}, обучаемых {trainable_params:,}")
    
    # DataLoader'ы
    train_dataset = TensorDataset(
        torch.FloatTensor(data['X_h1_train']),
        torch.FloatTensor(data['X_d1_train']),
        torch.FloatTensor(data['X_w1_train']),
        torch.LongTensor(data['y_train']),
    )
    val_dataset = TensorDataset(
        torch.FloatTensor(data['X_h1_val']),
        torch.FloatTensor(data['X_d1_val']),
        torch.FloatTensor(data['X_w1_val']),
        torch.LongTensor(data['y_val']),
    )
    
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, drop_last=True)
    val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)
    
    # Веса классов
    class_weights = compute_class_weights(data['y_train'], num_classes=num_classes)
    logger.info(f"  Веса классов: {class_weights.tolist()}")
    
    # Функция потерь
    if use_focal_loss:
        criterion = FocalLoss(alpha=class_weights, gamma=focal_gamma)
    else:
        criterion = nn.CrossEntropyLoss(weight=class_weights)
    
    # Оптимизатор
    optimizer = AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)
    
    # Планировщик
    from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts
    scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=10, T_mult=2)
    
    # История
    history = {'train_loss': [], 'val_loss': [], 'val_accuracy': [],
               'val_macro_f1': [], 'lr': []}
    best_val_loss = float('inf')
    best_state_dict = None
    best_epoch = 0
    early_stop_counter = 0
    
    for epoch in range(n_epochs):
        # Обучение
        model.train()
        train_loss = 0.0
        train_batches = 0
        
        for batch in train_loader:
            h1_x, d1_x, w1_x, y_batch = batch
            h1_x, d1_x, w1_x = h1_x.to(device), d1_x.to(device), w1_x.to(device)
            y_batch = y_batch.to(device)
            
            optimizer.zero_grad()
            logits = model(h1_x, d1_x, w1_x)
            loss = criterion(logits, y_batch)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
            optimizer.step()
            scheduler.step(epoch + train_batches / len(train_loader))
            
            train_loss += loss.item()
            train_batches += 1
        
        avg_train_loss = train_loss / max(train_batches, 1)
        current_lr = optimizer.param_groups[0]['lr']
        
        # Валидация (MTF: (h1, d1, w1, y), а не (x, y))
        model.eval()
        val_loss = 0.0
        all_preds, all_targets = [], []
        with torch.no_grad():
            for batch in val_loader:
                h1_x, d1_x, w1_x, y_batch = batch
                h1_x, d1_x, w1_x = h1_x.to(device), d1_x.to(device), w1_x.to(device)
                y_batch = y_batch.to(device)
                logits = model(h1_x, d1_x, w1_x)
                loss = criterion(logits, y_batch)
                val_loss += loss.item()
                preds = logits.argmax(dim=1)
                all_preds.append(preds.cpu())
                all_targets.append(y_batch.cpu())
        
        val_loss /= len(val_loader)
        all_preds = torch.cat(all_preds)
        all_targets = torch.cat(all_targets)
        val_accuracy = (all_preds == all_targets).float().mean().item()
        
        # Macro F1
        from sklearn.metrics import f1_score
        val_macro_f1 = f1_score(all_targets.numpy(), all_preds.numpy(), average='macro')
        
        history['train_loss'].append(avg_train_loss)
        history['val_loss'].append(val_loss)
        history['val_accuracy'].append(val_accuracy)
        history['val_macro_f1'].append(val_macro_f1)
        history['lr'].append(current_lr)
        
        logger.info(
            f"  Epoch {epoch+1:3d}/{n_epochs} | "
            f"Train Loss: {avg_train_loss:.4f} | "
            f"Val Loss: {val_loss:.4f} | "
            f"Val Acc: {val_accuracy:.4f} | "
            f"Macro F1: {val_macro_f1:.4f} | "
            f"ES: {early_stop_counter}/{patience}"
        )
        
        # Early stopping
        if val_loss < best_val_loss:
            best_val_loss = val_loss
            best_state_dict = model.state_dict().copy()
            best_epoch = epoch
            early_stop_counter = 0
        else:
            early_stop_counter += 1
            if early_stop_counter >= patience:
                logger.info(f"  Early stopping на эпохе {epoch+1}")
                break
    
    # Восстановление лучшей модели
    model.load_state_dict(best_state_dict)
    model.eval()
    
    logger.info(f"\nЛучшая эпоха: {best_epoch+1}, Val Loss: {best_val_loss:.4f}")
    
    # Сохранение
    save_path = Path(save_dir) / model_name
    save_path.mkdir(parents=True, exist_ok=True)
    
    torch.save({
        'epoch': best_epoch,
        'model_state_dict': best_state_dict,
        'optimizer_state_dict': optimizer.state_dict(),
        'val_loss': best_val_loss,
        'val_accuracy': history['val_accuracy'][best_epoch],
        'val_macro_f1': history['val_macro_f1'][best_epoch],
        'history': history,
        'model_type': 'mtf',
        'num_classes': num_classes,
        'input_size': n_features,
        'feature_names': data.get('feature_cols', []),
    }, str(save_path / 'best_model.pt'))
    
    logger.info(f"  Сохранено: {save_path / 'best_model.pt'}")
    
    # Тест
    test_dataset = TensorDataset(
        torch.FloatTensor(data['X_h1_test']),
        torch.FloatTensor(data['X_d1_test']),
        torch.FloatTensor(data['X_w1_test']),
        torch.LongTensor(data['y_test']),
    )
    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
    
    # Тест (MTF: свою валидацию)
    model.eval()
    test_loss = 0.0
    all_preds, all_targets = [], []
    with torch.no_grad():
        for batch in test_loader:
            h1_x, d1_x, w1_x, y_batch = batch
            h1_x, d1_x, w1_x = h1_x.to(device), d1_x.to(device), w1_x.to(device)
            y_batch = y_batch.to(device)
            logits = model(h1_x, d1_x, w1_x)
            loss = criterion(logits, y_batch)
            test_loss += loss.item()
            preds = logits.argmax(dim=1)
            all_preds.append(preds.cpu())
            all_targets.append(y_batch.cpu())
    
    test_loss /= len(test_loader)
    all_preds = torch.cat(all_preds)
    all_targets = torch.cat(all_targets)
    test_accuracy = (all_preds == all_targets).float().mean().item()
    test_macro_f1 = f1_score(all_targets.numpy(), all_preds.numpy(), average='macro')
    logger.info(f"\nТест: Loss={test_loss:.4f}, Acc={test_accuracy:.4f}, Macro F1={test_macro_f1:.4f}")
    
    metrics = {
        'val_loss': best_val_loss,
        'val_accuracy': history['val_accuracy'][best_epoch],
        'val_macro_f1': history['val_macro_f1'][best_epoch],
        'test_loss': test_loss,
        'test_accuracy': test_accuracy,
        'test_macro_f1': test_macro_f1,
        'train_loss_first': history['train_loss'][0],
        'train_loss_last': history['train_loss'][-1],
        'best_epoch': best_epoch,
        'n_epochs_trained': len(history['train_loss']),
    }
    
    return model, metrics


def main():
    parser = argparse.ArgumentParser(description='Обучение MTF Wyckoff-модели')
    parser.add_argument('--limit', type=int, default=2000, help='Свечей на тикер на ТФ')
    parser.add_argument('--epochs', type=int, default=100, help='Макс. эпох')
    parser.add_argument('--batch-size', type=int, default=32, help='Размер батча')
    parser.add_argument('--lr', type=float, default=1e-3, help='Learning rate')
    parser.add_argument('--patience', type=int, default=30, help='Early stopping')
    parser.add_argument('--dropout', type=float, default=0.4, help='Dropout')
    parser.add_argument('--seed', type=int, default=42, help='Seed')
    parser.add_argument('--no-save', action='store_true', help='Не сохранять')
    args = parser.parse_args()
    
    set_seed(args.seed)
    
    # 1. Загрузка данных по всем тикерам и ТФ
    logger.info(f"{'=' * 70}")
    logger.info("ЭТАП 1: ЗАГРУЗКА ДАННЫХ")
    logger.info(f"{'=' * 70}")
    
    all_data = load_all_tickers_mtf(TICKERS, limit=args.limit, use_simple_phases=True)
    
    # 2. Глобальные признаки
    logger.info(f"\n{'=' * 70}")
    logger.info("ЭТАП 2: ОБЩИЕ ПРИЗНАКИ")
    logger.info(f"{'=' * 70}")
    
    feature_cols = compute_global_features(all_data)
    
    # 3. Сбор MTF датасета
    logger.info(f"\n{'=' * 70}")
    logger.info("ЭТАП 3: СБОР MTF ДАТАСЕТА")
    logger.info(f"{'=' * 70}")
    
    tf_config = {'H1': 60, 'D1': 30, 'W1': 10}
    data = build_mtf_dataset(all_data, feature_cols, tf_config, use_simple_phases=True)
    
    # 4. Обучение
    logger.info(f"\n{'=' * 70}")
    logger.info("ЭТАП 4: ОБУЧЕНИЕ")
    logger.info(f"{'=' * 70}")
    
    model_name = 'wyckoff_mt_mtf_v1'
    
    model, metrics = train_mtf_model(
        data=data,
        n_epochs=args.epochs,
        learning_rate=args.lr,
        weight_decay=1e-3,
        batch_size=args.batch_size,
        patience=args.patience,
        dropout=args.dropout,
        model_name=model_name,
        use_focal_loss=True,
        focal_gamma=2.0,
    )
    
    # 5. Регистрация
    if not args.no_save:
        save_dir = 'src/ml/models/saved'
        from src.ml.train.train_wyckoff import register_wyckoff_model
        register_wyckoff_model(
            model_name=model_name,
            model_type='mtf',
            ticker='MULTI',
            tf='MTF',
            metrics=metrics,
            save_dir=save_dir,
            num_classes=5,
            seq_len=30,
        )
        logger.info(f"\n✅ Модель сохранена: {save_dir}/{model_name}/")
        logger.info(f"   Название: {model_name}")
    
    logger.info(f"\n{'=' * 70}")
    logger.info(f"РЕЗУЛЬТАТЫ MTF:")
    logger.info(f"  Val:   Acc={metrics['val_accuracy']:.2%}, F1={metrics['val_macro_f1']:.2%}")
    logger.info(f"  Test:  Acc={metrics['test_accuracy']:.2%}, F1={metrics['test_macro_f1']:.2%}")
    logger.info(f"{'=' * 70}")


if __name__ == '__main__':
    main()
