"""
Скрипт обучения Wyckoff-модели на нескольких тикерах MOEX.

Собирает данные со всех доступных тикеров MOEX, размечает фазы Вайкоффа,
и обучает нейросеть на объединённом датасете.

Использование:
    python src/ml/train/train_wyckoff_multiticker.py                         # все тикеры
    python src/ml/train/train_wyckoff_multiticker.py --tickers SBER GAZP X5  # выборочно
    python src/ml/train/train_wyckoff_multiticker.py --model transformer     # другая архитектура
"""

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

import numpy as np
import pandas as pd
import torch

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.data.dataset import create_sequences
from src.ml.train.train_wyckoff import (
    FocalLoss, train_wyckoff_model, register_wyckoff_model,
    PHASE_NAMES_SIMPLE, PHASE_NAMES_FULL,
)
from src.ml.models.wyckoff import create_wyckoff_model

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

# Все MOEX тикеры с D1 данными
DEFAULT_TICKERS = ['ASTR', 'GAZP', 'LKOH', 'MOEX', 'MTSS', 'NSVZ',
                   'NVTK', 'PHOR', 'PLZL', 'ROSN', 'SBER', 'SNGSP', 'VTBR', 'X5']


def load_and_label_ticker(
    ticker: str,
    tf: str = 'D1',
    limit: int = 2000,
    use_simple_phases: bool = True,
) -> Optional[pd.DataFrame]:
    """
    Загрузить данные тикера, добавить признаки и разметить фазы Вайкоффа.

    Args:
        ticker: тикер MOEX.
        tf: таймфрейм.
        limit: макс. свечей.
        use_simple_phases: упрощённые фазы (5 классов).

    Returns:
        DataFrame с признаками и целевой переменной или None (если данных мало).
    """
    try:
        df = fetch_ohlcv_combined(ticker, tf, limit=limit)
        if df is None or len(df) < 100:
            logger.warning(f"  {ticker}: недостаточно данных ({len(df) if df is not None else 0})")
            return None

        target_col = 'wyckoff_phase_simple' if use_simple_phases else 'wyckoff_phase'

        # Price-признаки
        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)

        # Wyckoff/VSA признаки
        df = add_wyckoff_features(df)

        # Разметка фаз
        labeler = WyckoffLabeler(smooth_window=5)
        df = labeler.label_phases(df)

        # Очистка NaN
        df = df.ffill().dropna(subset=[target_col])

        # Определяем feature_cols
        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)

        # Удаляем NaN в признаках
        nan_mask = df[feature_cols].isna().any(axis=1)
        df = df[~nan_mask].copy()

        if len(df) < 50:
            logger.warning(f"  {ticker}: всего {len(df)} строк после очистки — пропускаем")
            return None

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

        return df

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


def build_multiticker_dataset(
    tickers: List[str],
    tf: str = 'D1',
    limit: int = 2000,
    seq_len: int = 30,
    val_split: float = 0.15,
    test_split: float = 0.15,
    use_simple_phases: bool = True,
) -> Dict[str, Any]:
    """
    Собрать и подготовить мульти-тикерный датасет.

    Args:
        tickers: список тикеров.
        tf: таймфрейм.
        limit: макс. свечей на тикер.
        seq_len: длина последовательности.
        val_split: доля валидации.
        test_split: доля теста.
        use_simple_phases: упрощённые фазы.

    Returns:
        Словарь с данными: X_train, y_train, X_val, y_val, X_test, y_test,
        feature_names, scaler, metadata.
    """
    from sklearn.preprocessing import StandardScaler

    logger.info(f"{'=' * 70}")
    logger.info(f"СБОР МУЛЬТИ-ТИКЕРНОГО ДАТАСЕТА")
    logger.info(f"  Тикеров: {len(tickers)}")
    logger.info(f"  ТФ: {tf} | seq_len: {seq_len}")
    logger.info(f"{'=' * 70}")

    target_col = 'wyckoff_phase_simple' if use_simple_phases else 'wyckoff_phase'

    all_dfs = []
    global_feature_cols = None

    for ticker in tickers:
        logger.info(f"\nЗагрузка {ticker}...")
        df = load_and_label_ticker(ticker, tf, limit, use_simple_phases)
        if df is not None:
            # Нормализуем feature_cols (убеждаемся что одинаковые)
            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)

            if global_feature_cols is None:
                global_feature_cols = set(feature_cols)
            else:
                # Берём пересечение признаков (общие для всех тикеров)
                global_feature_cols = global_feature_cols.intersection(set(feature_cols))

            all_dfs.append(df)

    if not all_dfs:
        raise ValueError("Нет данных ни для одного тикера!")

    # Убеждаемся, что у всех тикеров одинаковые колонки
    global_feature_cols = sorted(global_feature_cols)
    logger.info(f"\nОбщих признаков: {len(global_feature_cols)}")

    # Собираем X и y
    X_list, y_list = [], []
    for df in all_dfs:
        # Берём только общие признаки (в том же порядке)
        X_list.append(df[global_feature_cols].values.astype(np.float32))
        y_list.append(df[target_col].values.astype(np.int64))

    # Объединяем
    X_all = np.vstack(X_list)
    y_all = np.concatenate(y_list)

    logger.info(f"Всего данных: {len(X_all)} строк, {len(global_feature_cols)} признаков")

    # Распределение фаз
    phase_dist = pd.Series(y_all).value_counts().sort_index()
    logger.info(f"Распределение фаз:")
    for pid in sorted(phase_dist.index):
        name = PHASE_NAMES_SIMPLE.get(pid, PHASE_NAMES_FULL.get(pid, f'C{pid}'))
        logger.info(f"  {pid}: {name} — {phase_dist[pid]} ({phase_dist[pid]/len(y_all)*100:.1f}%)")

    # Хронологическое разбиение (на уровне строк, не последовательностей)
    n = len(X_all)
    test_n = max(int(n * test_split), seq_len + 5)
    val_n = max(int(n * val_split), seq_len + 5)
    train_n = n - val_n - test_n

    if train_n < seq_len + 10:
        # Слишком мало train — корректируем
        val_n = seq_len + 5
        test_n = seq_len + 5
        train_n = n - val_n - test_n

    X_train_raw = X_all[:train_n]
    y_train_raw = y_all[:train_n]
    X_val_raw = X_all[train_n:train_n + val_n]
    y_val_raw = y_all[train_n:train_n + val_n]
    X_test_raw = X_all[train_n + val_n:]
    y_test_raw = y_all[train_n + val_n:]

    logger.info(f"Разбиение (raw): train={train_n}, val={len(X_val_raw)}, test={len(X_test_raw)}")

    # Нормализация
    scaler = StandardScaler()
    scaler.fit(X_train_raw)

    X_train_norm = scaler.transform(X_train_raw)
    X_val_norm = scaler.transform(X_val_raw)
    X_test_norm = scaler.transform(X_test_raw)

    # Создание последовательностей
    X_train, y_train = create_sequences(X_train_norm, y_train_raw, seq_len)
    X_val, y_val = create_sequences(X_val_norm, y_val_raw, seq_len)
    X_test, y_test = create_sequences(X_test_norm, y_test_raw, seq_len)

    logger.info(f"Sequences: train {X_train.shape}, val {X_val.shape}, test {X_test.shape}")

    # Распределение
    for name, yy in [('Train', y_train), ('Val', y_val), ('Test', y_test)]:
        classes, counts = np.unique(yy, return_counts=True)
        dist = {int(c): int(cnt) for c, cnt in zip(classes, counts)}
        logger.info(f"  {name}: {dist}")

    num_classes = 5 if use_simple_phases else 8

    metadata = {
        'tickers': tickers,
        'tf': tf,
        'seq_len': seq_len,
        'num_classes': num_classes,
        'use_simple_phases': use_simple_phases,
        'n_features': len(global_feature_cols),
        'total_raw_samples': len(X_all),
        'train_samples': len(X_train),
        'val_samples': len(X_val),
        'test_samples': len(X_test),
        'phase_distribution': {int(k): int(v) for k, v in phase_dist.items()},
    }

    return {
        'X_train': X_train, 'y_train': y_train,
        'X_val': X_val, 'y_val': y_val,
        'X_test': X_test, 'y_test': y_test,
        'feature_names': global_feature_cols,
        'scaler': scaler,
        'metadata': metadata,
        'num_classes': num_classes,
    }


def main():
    parser = argparse.ArgumentParser(
        description='Обучение Wyckoff-модели на нескольких тикерах MOEX'
    )
    parser.add_argument('--tickers', nargs='+', default=None,
                       help='Тикеры для обучения (по умолчанию все 14 MOEX)')
    parser.add_argument('--tf', type=str, default='D1',
                       help='Таймфрейм')
    parser.add_argument('--model', type=str, default='lstm',
                       choices=['lstm', 'cnn', 'transformer'],
                       help='Архитектура модели')
    parser.add_argument('--seq-len', type=int, default=30,
                       help='Длина последовательности')
    parser.add_argument('--epochs', type=int, default=100,
                       help='Максимальное количество эпох')
    parser.add_argument('--batch-size', type=int, default=64,
                       help='Размер батча')
    parser.add_argument('--lr', type=float, default=1e-3,
                       help='Learning rate')
    parser.add_argument('--patience', type=int, default=20,
                       help='Терпение для early stopping')
    parser.add_argument('--seed', type=int, default=42,
                       help='Seed')
    parser.add_argument('--limit', type=int, default=2000,
                       help='Свечей на тикер')
    parser.add_argument('--no-save', action='store_true',
                       help='Не сохранять модель')
    args = parser.parse_args()

    tickers = args.tickers or DEFAULT_TICKERS

    # Отключаем wandb/matplotlib если есть
    import warnings
    warnings.filterwarnings('ignore')

    from src.ml.train.trainer import set_seed
    set_seed(args.seed)

    model_name = f'wyckoff_mt_{args.tf.lower()}_{args.model}_v1'
    save_dir = 'src/ml/models/saved'

    # 1. Сбор мульти-тикерного датасета
    data = build_multiticker_dataset(
        tickers=tickers,
        tf=args.tf,
        limit=args.limit,
        seq_len=args.seq_len,
        use_simple_phases=True,
    )

    # 2. Обучение модели
    model_params = {}
    if args.model == 'lstm':
        model_params = {'hidden_size': 128, 'num_layers': 2, 'dropout': 0.4}
    elif args.model == 'cnn':
        model_params = {'dropout': 0.3}
    elif args.model == 'transformer':
        model_params = {'d_model': 128, 'nhead': 4, 'num_layers': 3,
                        'dim_feedforward': 256, 'dropout': 0.3}

    model, train_result = train_wyckoff_model(
        data=data,
        model_type=args.model,
        model_params=model_params,
        n_epochs=args.epochs,
        learning_rate=args.lr,
        batch_size=args.batch_size,
        patience=args.patience,
        model_name=model_name,
        save_dir=save_dir,
        use_focal_loss=True,
        focal_gamma=2.0,
        use_simple_phases=True,
    )

    # 3. Регистрация модели
    if not args.no_save:
        register_wyckoff_model(
            model_name=model_name,
            model_type=args.model,
            ticker='MULTI',
            tf=args.tf,
            metrics=train_result,
            save_dir=save_dir,
            num_classes=data['num_classes'],
            seq_len=args.seq_len,
        )
        logger.info(f"\n✅ Модель сохранена: {save_dir}/{model_name}/")
        logger.info(f"   Название: {model_name}")

    logger.info(f"\n{'=' * 70}")
    logger.info(f"ОБУЧЕНИЕ ЗАВЕРШЕНО")
    logger.info(f"{'=' * 70}")


if __name__ == '__main__':
    main()
