"""
Модуль инференса для определения фаз рынка по методу Вайкоффа.

Позволяет загрузить обученную модель и определить текущую фазу рынка
для любого тикера MOEX на основе OHLCV данных.

Usage:
    from src.ml.inference.wyckoff_inference import WyckoffInference

    inferer = WyckoffInference()
    result = inferer.analyze('SBER', tf='D1')
    print(result['phase_name'], result['confidence'])

    # Пакетный анализ нескольких тикеров
    results = inferer.batch_analyze(['SBER', 'GAZP', 'X5'])
"""

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

import numpy as np
import pandas as pd
import torch
import torch.nn as nn

from src.ml.models.registry import ModelRegistry

# Настройка логирования
logger = logging.getLogger('wyckoff_inference')


# Константы фаз
PHASE_NAMES_FULL = {
    0: 'Сильный Маркдаун (Strong Markdown)',
    1: 'Маркдаун (Markdown)',
    2: 'Накопление Phase A (SC/AR/ST)',
    3: 'Накопление Phase B-C (Spring/LPS)',
    4: 'Накопление Phase D-E / Начало Маркапа',
    5: 'Маркап (Markup)',
    6: 'Распределение Phase A-B (BC/AR/ST)',
    7: 'Распределение Phase C-D (UT/LPSY)',
}

PHASE_NAMES_SIMPLE = {
    0: 'Маркдаун (Markdown)',
    1: 'Накопление (раннее) / Accumulation (early)',
    2: 'Накопление (позднее) / Начало Маркапа',
    3: 'Маркап (Markup)',
    4: 'Распределение (Distribution)',
}

# Цветовые коды для фаз
PHASE_COLORS = {
    0: '\033[91m',   # Красный — сильный маркдаун
    1: '\033[91m',   # Красный — маркдаун
    2: '\033[93m',   # Жёлтый — накопление A
    3: '\033[93m',   # Жёлтый — накопление BC
    4: '\033[92m',   # Зелёный — накопление DE / начало маркапа
    5: '\033[92m',   # Зелёный — маркап
    6: '\033[95m',   # Пурпурный — распределение AB
    7: '\033[95m',   # Пурпурный — распределение CD
}
RESET_COLOR = '\033[0m'


class WyckoffPhaseResult:
    """
    Результат определения фазы Вайкоффа.

    Attributes:
        ticker: тикер инструмента.
        phase: числовой идентификатор фазы.
        phase_name: человекочитаемое название фазы.
        confidence: уверенность модели (0-100%).
        phase_probs: вероятности всех фаз.
        events: обнаруженные VSA/Wyckoff события.
        is_accumulation: True если фаза накопления.
        is_distribution: True если фаза распределения.
        is_markup: True если фаза маркапа.
        is_markdown: True если фаза маркдауна.
    """

    def __init__(
        self,
        ticker: str,
        phase: int,
        confidence: float,
        phase_probs: Optional[Dict[int, float]] = None,
        events: Optional[Dict[str, float]] = None,
        is_simple: bool = False,
    ):
        self.ticker = ticker
        self.phase = phase
        self.confidence = confidence
        self.phase_probs = phase_probs or {}
        self.events = events or {}
        self.is_simple = is_simple

        self.phase_names = PHASE_NAMES_SIMPLE if is_simple else PHASE_NAMES_FULL
        self.phase_name = self.phase_names.get(phase, f'Unknown ({phase})')

        # Тип фазы
        if is_simple:
            self.is_accumulation = phase in [1, 2]
            self.is_markup = phase == 3
            self.is_distribution = phase == 4
            self.is_markdown = phase == 0
        else:
            self.is_accumulation = phase in [2, 3, 4]
            self.is_markup = phase == 5
            self.is_distribution = phase in [6, 7]
            self.is_markdown = phase in [0, 1]

    def __str__(self) -> str:
        color = PHASE_COLORS.get(self.phase, '')
        s = (
            f"{self.ticker:6s} | "
            f"{color}{self.phase_name:40s}{RESET_COLOR} | "
            f"conf: {self.confidence:5.1f}%"
        )
        if self.events:
            active_events = [k for k, v in self.events.items() if v > 0]
            if active_events:
                s += f" | events: {', '.join(active_events[:3])}"
        return s

    def to_dict(self) -> Dict[str, Any]:
        """Конвертировать в словарь для JSON/отчёта."""
        return {
            'ticker': self.ticker,
            'phase': int(self.phase),
            'phase_name': self.phase_name,
            'confidence': round(float(self.confidence), 1),
            'phase_probs': {str(k): round(float(v), 3) for k, v in self.phase_probs.items()},
            'events': {str(k): round(float(v), 3) for k, v in self.events.items()},
            'is_accumulation': self.is_accumulation,
            'is_markup': self.is_markup,
            'is_distribution': self.is_distribution,
            'is_markdown': self.is_markdown,
        }


class WyckoffInference:
    """
    Класс для инференса Wyckoff-модели.

    Позволяет загрузить обученную модель и определять фазу Вайкоффа
    для любого тикера в реальном времени.

    Args:
        model_name: имя модели в ModelRegistry (если None — ищет лучшую).
        ticker: тикер, для которого загружать модель (если model_name=None).
        tf: таймфрейм.
        device: устройство ('cpu' или 'cuda').
        use_rule_based_fallback: использовать rule-based как fallback.
    """

    def __init__(
        self,
        model_name: Optional[str] = None,
        ticker: str = 'X5',
        tf: str = 'D1',
        device: str = 'cpu',
        use_rule_based_fallback: bool = True,
    ):
        self.ticker = ticker
        self.tf = tf
        self.device = torch.device(device if torch.cuda.is_available() else 'cpu')
        self.use_rule_based_fallback = use_rule_based_fallback

        # Загрузка модели
        self.model = None
        self.metadata = None
        self.scaler = None
        self.feature_names = None
        self.model_loaded = self._load_model(model_name, ticker, tf)

        # Labeler для fallback
        self.labeler = None
        if use_rule_based_fallback:
            from src.ml.data.wyckoff_labeling import WyckoffLabeler
            self.labeler = WyckoffLabeler()

        if self.model_loaded:
            logger.info(
                f"✅ Wyckoff-модель загружена: {model_name or 'auto'}\n"
                f"   Трейнер: {ticker} | ТФ: {tf}\n"
                f"   Классов: {self.metadata.get('num_classes', '?')}\n"
                f"   Seq len: {self.metadata.get('seq_len', '?')}\n"
                f"   Устройство: {self.device}"
            )

    def _load_model(
        self,
        model_name: Optional[str],
        ticker: str,
        tf: str,
    ) -> bool:
        """
        Загрузить модель из ModelRegistry или файла.

        Args:
            model_name: имя модели.
            ticker: тикер.
            tf: таймфрейм.

        Returns:
            True если модель загружена успешно.
        """
        registry = ModelRegistry()

        # Поиск модели
        model_entry = None
        if model_name:
            model_entry = registry.get_model(model_name)
        else:
            # Ищем любую Wyckoff-модель (мульти-тикерную или для тикера)
            all_models = registry.list_models(status='active')
            wyckoff_models = [
                m for m in all_models
                if m.get('params', {}).get('purpose') == 'wyckoff_phase_detection'
                and m.get('timeframe') == tf
            ]
            if wyckoff_models:
                # Берём лучшую по val_macro_f1
                model_entry = max(
                    wyckoff_models,
                    key=lambda m: m.get('metrics', {}).get('val_macro_f1', 0)
                )
                if model_entry:
                    logger.info(f"  Выбрана модель: {model_entry['model_name']} "
                                f"(F1={model_entry.get('metrics', {}).get('val_macro_f1', '?'):.2%})")

        if model_entry is None:
            logger.warning(
                f"Wyckoff-модель не найдена для {ticker} {tf}. "
                f"Будет использован rule-based метод."
            )
            return False

        self.metadata = {
            'model_name': model_entry['model_name'],
            'model_type': model_entry['model_type'],
            'num_classes': model_entry.get('params', {}).get('num_classes', 8),
            'seq_len': model_entry.get('params', {}).get('seq_len', 60),
            'metrics': model_entry.get('metrics', {}),
        }

        # Загрузка файла модели
        model_path = Path(model_entry['model_path'])
        model_file = model_path / 'best_model.pt'

        if not model_file.exists():
            # Пробуем другие варианты
            for fname in ['wyckoff_final.pt', 'model.pt', 'model_checkpoint.pt']:
                alt_path = model_path / fname
                if alt_path.exists():
                    model_file = alt_path
                    break

        if not model_file.exists():
            logger.error(f"Файл модели не найден: {model_path}")
            return False

        # Загружаем чекпоинт для получения input_size
        checkpoint = torch.load(str(model_file), map_location='cpu', weights_only=False)
        if 'model_state_dict' in checkpoint:
            state_dict = checkpoint['model_state_dict']
        else:
            state_dict = checkpoint

        # Получаем input_size из чекпоинта (сохранён в обучении)
        input_size = checkpoint.get('input_size', None)
        if input_size is None:
            # Определяем из формы весов первого слоя
            for key in state_dict:
                if 'lstm.weight_ih_l0' in key or 'input_proj.weight' in key or 'input_norm.weight' in key:
                    weight_shape = state_dict[key].shape
                    if len(weight_shape) == 2:
                        input_size = weight_shape[1]  # [out_features, in_features]
                    elif len(weight_shape) == 1:
                        input_size = weight_shape[0]
                    break

        # Сохраняем feature_names из чекпоинта (для правильного выбора признаков)
        self.metadata['feature_names'] = checkpoint.get('feature_names', [])

        arch = model_entry.get('params', {}).get('arch', 'lstm')
        num_classes = self.metadata['num_classes']
        seq_len = self.metadata['seq_len']

        if input_size is None:
            logger.error("Не удалось определить input_size из чекпоинта")
            return False

        from src.ml.models.wyckoff import create_wyckoff_model
        self.model = create_wyckoff_model(
            model_type=arch,
            input_size=input_size,
            num_classes=num_classes,
            seq_len=seq_len,
        )

        # Загрузка весов
        try:
            self.model.load_state_dict(state_dict)
            self.model.eval()
            self.model.to(self.device)
            self.metadata['input_size'] = input_size
            logger.info(f"  input_size={input_size}, num_classes={num_classes}, seq_len={seq_len}")
            return True

        except Exception as e:
            logger.error(f"Ошибка загрузки весов: {e}")
            return False

    def _prepare_data(
        self,
        df: pd.DataFrame,
    ) -> Optional[np.ndarray]:
        """
        Подготовить данные для инференса.

        Args:
            df: DataFrame с OHLCV данными.

        Returns:
            Нормализованная последовательность для модели или None.
        """
        try:
            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

            # Добавление признаков + разметка Вайкоффа (для генерации событий)
            df = add_price_features(df)
            df = add_indicator_features(df)
            df = add_wyckoff_features(df)

            # Rule-based разметка — генерирует события (spring, sos, ar, st, lps, lpsy и т.д.)
            # Это необходимо, т.к. модель обучена на фичах, включающих эти колонки
            labeler = WyckoffLabeler(smooth_window=5)
            df = labeler.label_phases(df)

            # Выбор признаков — используем сохранённый список из чекпоинта
            train_features = self.metadata.get('feature_names', [])
            if train_features:
                feature_cols = [c for c in train_features if c in df.columns]
                missing = set(train_features) - set(feature_cols)
                if missing:
                    logger.warning(f"Отсутствуют признаки: {missing}")
                    for c in missing:
                        df[c] = 0.0
                    feature_cols = list(train_features)
            else:
                # Fallback: ручной отбор
                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]]

            # Обработка NaN
            df = df.ffill().bfill().fillna(0)

            X = df[feature_cols].values.astype(np.float32)

            # Нормализация (Online — используем статистики последних N свечей)
            mean = X[-200:].mean(axis=0)
            std = X[-200:].std(axis=0) + 1e-8
            X_norm = (X - mean) / std

            # Берём последние seq_len свечей
            seq_len = self.metadata.get('seq_len', 60)
            if len(X_norm) < seq_len:
                logger.error(f"Недостаточно данных: {len(X_norm)} < {seq_len}")
                return None

            X_seq = X_norm[-seq_len:]
            X_seq = np.expand_dims(X_seq, axis=0)  # (1, seq_len, n_features)

            return X_seq.astype(np.float32)

        except Exception as e:
            logger.error(f"Ошибка подготовки данных: {e}")
            return None

    def _predict_from_model(
        self,
        df: pd.DataFrame,
    ) -> Optional[WyckoffPhaseResult]:
        """
        Получить предсказание от нейросетевой модели.

        Args:
            df: DataFrame с OHLCV.

        Returns:
            WyckoffPhaseResult или None.
        """
        if self.model is None:
            return None

        X_seq = self._prepare_data(df)
        if X_seq is None:
            return None

        with torch.no_grad():
            x_tensor = torch.FloatTensor(X_seq).to(self.device)
            logits = self.model(x_tensor)
            probs = torch.softmax(logits, dim=1)

        probs_np = probs.cpu().numpy()[0]
        phase = int(np.argmax(probs_np))
        confidence = float(probs_np[phase] * 100)

        phase_probs = {i: float(p) for i, p in enumerate(probs_np)}

        # Определяем события из последних свечей
        events = self._detect_events(df)

        num_classes = self.metadata.get('num_classes', 8)
        is_simple = num_classes == 5

        return WyckoffPhaseResult(
            ticker=self.ticker,
            phase=phase,
            confidence=confidence,
            phase_probs=phase_probs,
            events=events,
            is_simple=is_simple,
        )

    def _detect_events(self, df: pd.DataFrame) -> Dict[str, float]:
        """
        Определить VSA/Wyckoff события в последних свечах.

        Args:
            df: DataFrame с VSA-метриками.

        Returns:
            Словарь событий с их интенсивностью.
        """
        events = {}
        if df.empty:
            return events

        last_rows = df.tail(10)

        # VSA сигналы (если есть)
        for signal in ['stopping_volume', 'selling_climax', 'buying_climax',
                       'no_demand', 'no_supply', 'upthrust', 'test_volume']:
            if signal in last_rows.columns:
                count = int(last_rows[signal].sum())
                if count > 0:
                    events[signal] = count / len(last_rows)

        # Wyckoff события (если есть)
        for event in ['spring', 'shakeout', 'sos', 'lps',
                      'upthrust_detected', 'lpsy', 'ar', 'st']:
            if event in last_rows.columns:
                count = int(last_rows[event].sum())
                if count > 0:
                    events[event] = count / len(last_rows)

        return events

    def _predict_rule_based(
        self,
        df: pd.DataFrame,
    ) -> WyckoffPhaseResult:
        """
        Rule-based определение фазы Вайкоффа (без нейросети).

        Args:
            df: DataFrame с OHLCV.

        Returns:
            WyckoffPhaseResult.
        """
        try:
            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, add_vsa_metrics, add_vsa_signals

            df = add_price_features(df)
            df = add_indicator_features(df)
            df = add_wyckoff_features(df)

            if self.labeler is None:
                from src.ml.data.wyckoff_labeling import WyckoffLabeler
                self.labeler = WyckoffLabeler()

            df = self.labeler.label_phases(df)

            if 'wyckoff_phase_simple' in df.columns:
                phase = int(df['wyckoff_phase_simple'].iloc[-1])
                is_simple = True
                phase_name_col = 'wyckoff_phase_simple_name'
            else:
                phase = int(df['wyckoff_phase'].iloc[-1])
                is_simple = False
                phase_name_col = 'wyckoff_phase_name'

            events = self._detect_events(df)

            # Rule-based confidence (основан на количестве подтверждающих событий)
            event_count = sum(v for v in events.values())
            confidence = min(50 + event_count * 10, 85)

            return WyckoffPhaseResult(
                ticker=self.ticker,
                phase=phase,
                confidence=confidence,
                events=events,
                is_simple=is_simple,
            )

        except Exception as e:
            logger.error(f"Rule-based prediction error: {e}")
            return WyckoffPhaseResult(
                ticker=self.ticker,
                phase=1,  # Markdown по умолчанию
                confidence=30,
                is_simple=True,
            )

    def analyze(
        self,
        ticker: Optional[str] = None,
        tf: Optional[str] = None,
        df: Optional[pd.DataFrame] = None,
    ) -> WyckoffPhaseResult:
        """
        Определить фазу Вайкоффа для тикера.

        Args:
            ticker: тикер MOEX. Если None, используется тикер из конструктора.
            tf: таймфрейм. Если None, используется tf из конструктора.
            df: опциональный DataFrame с OHLCV (если не указан, загружается из БД).

        Returns:
            WyckoffPhaseResult с результатом анализа.
        """
        ticker = ticker or self.ticker
        tf = tf or self.tf
        self.ticker = ticker
        self.tf = tf

        # Если передан DataFrame, используем его
        if df is None:
            from src.db.connection import fetch_ohlcv_combined
            limit = 500
            df = fetch_ohlcv_combined(ticker, tf, limit=limit)
            if df is None or len(df) < 60:
                logger.error(f"Не удалось загрузить данные для {ticker} {tf}")
                return WyckoffPhaseResult(ticker, 1, 20)

        # Сначала пробуем нейросетевую модель
        result = None
        if self.model_loaded and self.model is not None:
            result = self._predict_from_model(df)

        # Fallback на rule-based
        if result is None and self.use_rule_based_fallback:
            logger.info(f"Использую rule-based метод для {ticker}")
            result = self._predict_rule_based(df)

        if result is None:
            result = WyckoffPhaseResult(ticker, 1, 20)

        return result

    def batch_analyze(
        self,
        tickers: List[str],
        tf: str = 'D1',
        verbose: bool = True,
    ) -> List[WyckoffPhaseResult]:
        """
        Пакетный анализ нескольких тикеров.

        Args:
            tickers: список тикеров.
            tf: таймфрейм.
            verbose: выводить результаты.

        Returns:
            Список WyckoffPhaseResult.
        """
        results = []
        for ticker in tickers:
            if verbose:
                logger.info(f"Анализ {ticker}...")
            result = self.analyze(ticker=ticker, tf=tf)
            results.append(result)

            if verbose:
                print(f"  {result}")

        return results

    def print_summary(self, results: List[WyckoffPhaseResult]) -> None:
        """Вывести сводку по результатам анализа."""
        print(f"\n{'=' * 80}")
        print(f"СВОДКА ФАЗ ВАЙКОФФА")
        print(f"{'=' * 80}")
        print(f"{'Тикер':<8} {'Фаза':<42} {'Conf':<8} {'Тип':<15}")
        print(f"{'-' * 80}")

        for r in results:
            color = PHASE_COLORS.get(r.phase, '')
            if r.is_accumulation:
                phase_type = '🟢 Накопление'
            elif r.is_markup:
                phase_type = '🟢 Маркап'
            elif r.is_distribution:
                phase_type = '🔴 Распределение'
            else:
                phase_type = '🔴 Маркдаун'

            print(
                f"{r.ticker:<8} "
                f"{color}{r.phase_name:<42}{RESET_COLOR} "
                f"{r.confidence:>5.1f}%  "
                f"{phase_type:<15}"
            )

        # Статистика
        phases = [r.phase for r in results]
        from collections import Counter
        phase_counts = Counter(phases)

        print(f"\n📊 Распределение:")
        for phase_id, count in sorted(phase_counts.items()):
            name = PHASE_NAMES_SIMPLE.get(phase_id, PHASE_NAMES_FULL.get(phase_id, f'Phase {phase_id}'))
            bar = '█' * count
            print(f"  {name:40s}: {bar} {count}")

        print(f"{'=' * 80}\n")


# CLI для прямого использования
if __name__ == '__main__':
    import argparse

    parser = argparse.ArgumentParser(description='Определение фаз Вайкоффа')
    parser.add_argument('tickers', nargs='+', help='Тикеры для анализа')
    parser.add_argument('--tf', default='D1', help='Таймфрейм')
    parser.add_argument('--model', default=None, help='Имя модели в ModelRegistry')
    parser.add_argument('--device', default='cpu', help='Устройство')
    parser.add_argument('--rule-only', action='store_true', help='Только rule-based')
    args = parser.parse_args()

    inferer = WyckoffInference(
        model_name=args.model,
        ticker=args.tickers[0] if args.tickers else 'X5',
        tf=args.tf,
        device=args.device,
        use_rule_based_fallback=args.rule_only or True,
    )

    results = inferer.batch_analyze(args.tickers, tf=args.tf)
    inferer.print_summary(results)
