"""
Сравнение определения фаз Вайкоффа: rule-based (tech_analysis.wyckoff_phase) vs NN ensemble (H1+D1+W1).

Запуск:
    python _compare_wyckoff_approaches.py

Результат:
    - Таблица сравнения по каждому тикеру
    - Матрица соответствия (confusion-like)
    - Метрики: Cohen's Kappa, Accuracy согласия, % совпадений
"""

import sys
import json
import warnings
from pathlib import Path
from collections import Counter

warnings.filterwarnings("ignore")
sys.path.insert(0, str(Path(__file__).resolve().parent))

import pandas as pd
import numpy as np
from sklearn.metrics import cohen_kappa_score, accuracy_score, confusion_matrix

# ── 1. Rule-based Wyckoff (из _analyze_all_db) ──────────────────────────
from src.analysis.tech_analysis import wyckoff_phase
from src.db.connection import fetch_ohlcv_combined
from src.indicators.calculations import calc_all_indicators

# ── 2. NN Ensemble ──────────────────────────────────────────────────────
from src.ml.inference.wyckoff_ensemble import WyckoffEnsemble

# ── MOEX tickers for this test ──────────────────────────────────────────
MOEX_TICKERS = [
    'SBER', 'GAZP', 'LKOH', 'ROSN', 'NVTK', 'MGNT', 'TATN', 'SNGS',
    'PLZL', 'PHOR', 'NLMK', 'CHMF', 'GMKN', 'ALRS', 'MTSS', 'VTBR',
    'MOEX', 'X5', 'AFLT', 'RUAL', 'SELG', 'IRAO', 'HYDR', 'MAGN',
    'RASP', 'NMTP', 'FESH', 'CBOM', 'TCSG', 'VKCO', 'YDEX', 'OZON',
    'ASTR',
]

# ── Фазы rule-based → numeric mapping ──────────────────────────────────
def map_rule_phase_to_class(phase_text: str) -> int:
    """Сопоставить текстовую фазу из rule-based с 5-классовой системой NN."""
    phase_lower = phase_text.lower()

    if 'маркдаун' in phase_lower or 'markdown' in phase_lower:
        return 0
    if 'накопление' in phase_lower or 'accumulation' in phase_lower:
        return 1
    if 'маркап' in phase_lower or 'markup' in phase_lower:
        return 3
    if 'распределение' in phase_lower or 'distribution' in phase_lower:
        return 4
    # Переходные / неопределённые фазы
    if 'консолидация' in phase_lower:
        return 2
    return 2  # default → neutral


PHASE_NAMES = ['Маркдаун', 'Накопление(раннее)', 'Накопление(позд)/Маркап', 'Маркап', 'Распределение']


def run_rule_based(ticker: str) -> dict:
    """Запустить rule-based wyckoff_phase (как в _analyze_all_db.py)."""
    try:
        df_w1 = fetch_ohlcv_combined(ticker, 'W1', limit=100)
        if df_w1 is None or len(df_w1) < 30:
            return {'phase': -1, 'phase_text': '—', 'class': -1, 'events': [], 'direction': '—'}

        df_w1 = calc_all_indicators(df_w1)
        result = wyckoff_phase(df_w1)
        phase_text = result.get('phase', '—')
        return {
            'phase': map_rule_phase_to_class(phase_text),
            'phase_text': phase_text,
            'class': map_rule_phase_to_class(phase_text),
            'events': result.get('events', []),
            'direction': result.get('direction', '—'),
        }
    except Exception as e:
        return {'phase': -1, 'phase_text': f'ERROR: {e}', 'class': -1, 'events': [], 'direction': '—'}


def run_nn_ensemble(ticker: str, ensemble: WyckoffEnsemble) -> dict:
    """Запустить NN ensemble для тикера."""
    try:
        result = ensemble.analyze(ticker)
        return {
            'phase': result.phase,
            'phase_text': result.phase_name,
            'class': result.phase,
            'confidence': result.confidence,
            'probs': result.phase_probs,
            'details': result.details,
        }
    except Exception as e:
        return {'phase': -1, 'phase_text': f'ERROR: {e}', 'class': -1, 'confidence': 0, 'probs': {}, 'details': {}}


def main():
    print("=" * 90)
    print("      СРАВНЕНИЕ: Rule-based (W1) vs NN Ensemble (H1+D1+W1)")
    print("=" * 90)

    # Загружаем ансамбль
    print("\n[1/3] Загрузка NN ensemble...")
    ensemble = WyckoffEnsemble(use_rule_based_fallback=True)
    print(f"  Загружено ТФ: {list(ensemble.inferers.keys())}")

    # Прогоняем оба подхода
    print(f"\n[2/3] Анализ {len(MOEX_TICKERS)} тикеров...")
    results = []
    for ticker in MOEX_TICKERS:
        rule = run_rule_based(ticker)
        nn = run_nn_ensemble(ticker, ensemble)
        results.append({
            'ticker': ticker,
            'rule_phase': rule['phase'],
            'rule_text': rule['phase_text'][:60],
            'rule_events': rule['events'],
            'nn_phase': nn['phase'],
            'nn_text': nn['phase_text'],
            'nn_conf': nn.get('confidence', 0),
            'nn_details': nn.get('details', {}),
            'agreement': rule['phase'] == nn['phase'],
        })

    # ── Вывод таблицы ──
    print("\n[3/3] Результаты:\n")
    print(f"{'Тикер':<7} {'Rule-based W1':<42} {'→cls':<5} {'NN Ensemble (H1+D1+W1)':<42} {'→cls':<5} {'Conf':<6} {'Совпадение':<10}")
    print("-" * 120)

    agree_count = 0
    total = 0

    for r in results:
        if r['rule_phase'] >= 0 and r['nn_phase'] >= 0:
            total += 1
            if r['agreement']:
                agree_count += 1

        rule_cls = str(r['rule_phase']) if r['rule_phase'] >= 0 else '?'
        nn_cls = str(r['nn_phase']) if r['nn_phase'] >= 0 else '?'

        mark = '✓' if r['agreement'] else '✗' if r['rule_phase'] >= 0 and r['nn_phase'] >= 0 else '?'

        # Детали по ТФ
        details_str = ''
        if 'nn_details' in r:
            for tf, d in r['nn_details'].items():
                details_str += f" {tf}:{d['phase_name'][:15]}({d['confidence']:.0f}%)"

        print(
            f"{r['ticker']:<7} "
            f"{r['rule_text']:<42} "
            f"{rule_cls:<5} "
            f"{r['nn_text']:<42} "
            f"{nn_cls:<5} "
            f"{r.get('nn_conf', 0):>5.1f}% "
            f"{mark:<10}"
        )

    # ── Матрица соответствия ──
    print("\n" + "=" * 90)
    print("      МАТРИЦА СООТВЕТСТВИЯ (Rule-based vs NN)")
    print("=" * 90)

    rule_classes = [r['rule_phase'] for r in results if r['rule_phase'] >= 0 and r['nn_phase'] >= 0]
    nn_classes = [r['nn_phase'] for r in results if r['rule_phase'] >= 0 and r['nn_phase'] >= 0]

    if len(rule_classes) > 0:
        cm = confusion_matrix(rule_classes, nn_classes, labels=[0, 1, 2, 3, 4])
        print(f"\n{'':>22} {'NN prediction':>50}")
        print(f"{'Rule \\ NN':<12}", end='')
        for i in range(5):
            print(f"{PHASE_NAMES[i][:12]:>12}", end='')
        print(f"{'Total':>8}")
        print("    " + "-" * 73)

        for i in range(5):
            row = cm[i] if i < len(cm) else [0]*5
            row_sum = sum(row)
            print(f"{PHASE_NAMES[i][:12]:<12}", end='')
            for j in range(5):
                val = row[j] if j < len(row) else 0
                if i == j:
                    print(f"\033[92m{val:>8}\033[0m", end='  ')
                elif val > 0:
                    print(f"\033[93m{val:>8}\033[0m", end='  ')
                else:
                    print(f"{val:>8}", end='  ')
            print(f"{row_sum:>6}")

    # ── Метрики ──
    print("\n" + "=" * 90)
    print("      МЕТРИКИ СОГЛАСИЯ")
    print("=" * 90)

    if len(rule_classes) > 0:
        kappa = cohen_kappa_score(rule_classes, nn_classes)
        acc = accuracy_score(rule_classes, nn_classes)
        total_valid = len(rule_classes)
        print(f"\n  Всего тикеров (с данными): {total_valid}")
        print(f"  Точное совпадение фаз:    {agree_count}/{total} = {agree_count/total*100:.1f}%")
        print(f"  Accuracy (согласие):       {acc*100:.1f}%")
        print(f"  Cohen's Kappa:             {kappa:.4f}")
        print(f"  Случайное согласие K=0     {'⬆ Выше случайного' if kappa > 0.1 else '⬇ На уровне случайного' if kappa > -0.1 else '⬇ Ниже случайного'}")

        if kappa >= 0.61:
            print(f"  Интерпретация Kappa:      \033[92mЗначительное согласие\033[0m")
        elif kappa >= 0.41:
            print(f"  Интерпретация Kappa:      \033[93mУмеренное согласие\033[0m")
        elif kappa >= 0.21:
            print(f"  Интерпретация Kappa:      \033[94mСлабое согласие\033[0m")
        else:
            print(f"  Интерпретация Kappa:      \033[91mМинимальное/отсутствует\033[0m")

    # ── Распределение фаз ──
    print("\n" + "=" * 90)
    print("      РАСПРЕДЕЛЕНИЕ ФАЗ")
    print("=" * 90)

    rule_dist = Counter(r['rule_phase'] for r in results if r['rule_phase'] >= 0)
    nn_dist = Counter(r['nn_phase'] for r in results if r['nn_phase'] >= 0)

    print(f"\n{'Фаза':<30} {'Rule-based':<15} {'NN Ensemble':<15}")
    print("-" * 60)
    for i in range(5):
        print(f"{PHASE_NAMES[i]:<30} {rule_dist.get(i, 0):<15} {nn_dist.get(i, 0):<15}")

    # ── Детальные расхождения ──
    print("\n" + "=" * 90)
    print("      ТИКЕРЫ С РАСХОЖДЕНИЕМ (>1 класс)")
    print("=" * 90)

    found = False
    for r in results:
        if r['rule_phase'] >= 0 and r['nn_phase'] >= 0:
            diff = abs(r['rule_phase'] - r['nn_phase'])
            if diff >= 2:
                found = True
                print(f"\n  {r['ticker']}: \033[91mRule={r['rule_text']} ({r['rule_phase']}) vs NN={r['nn_text']} ({r['nn_phase']})\033[0m")
                # Показать детали NN
                if 'nn_details' in r:
                    for tf, d in r['nn_details'].items():
                        print(f"           {tf}: {d['phase_name']} ({d['confidence']:.0f}%)")
                # Показать rule-based события
                if r['rule_events']:
                    for ev in r['rule_events']:
                        print(f"           Rule event: {ev}")
    if not found:
        print("  Нет расхождений >1 класса.")

    # ── Сохраняем полный JSON ──
    output_path = Path(__file__).resolve().parent / 'reports' / 'wyckoff_comparison.json'
    output_path.parent.mkdir(exist_ok=True)
    with open(output_path, 'w', encoding='utf-8') as f:
        json.dump({
            'comparison_results': [
                {
                    'ticker': r['ticker'],
                    'rule_phase': r['rule_phase'],
                    'rule_text': r['rule_text'],
                    'rule_events': r['rule_events'],
                    'rule_direction': r.get('direction', ''),
                    'nn_phase': r['nn_phase'],
                    'nn_text': r['nn_text'],
                    'nn_confidence': round(r.get('nn_conf', 0), 1),
                    'agreement': r['agreement'],
                }
                for r in results
            ],
            'metrics': {
                'accuracy': round(acc, 4) if len(rule_classes) > 0 else 0,
                'cohen_kappa': round(kappa, 4) if len(rule_classes) > 0 else 0,
                'total_valid': total_valid if len(rule_classes) > 0 else 0,
                'exact_match_count': agree_count,
                'exact_match_pct': round(agree_count/total*100, 1) if total > 0 else 0,
            },
            'confusion_matrix': cm.tolist() if len(rule_classes) > 0 else [],
        }, f, indent=2, ensure_ascii=False)

    print(f"\n📁 Полные результаты сохранены: {output_path}")
    print("\n" + "=" * 90)


if __name__ == '__main__':
    main()
