"""
Метрики производительности для классификации направлений BUY/HOLD/SELL.

Функции:
    calculate_accuracy        — точность классификации
    calculate_precision_recall_f1 — precision, recall, F1 по классам
    calculate_confusion_matrix — матрица ошибок
    classification_report     — полный отчёт по метрикам
"""

import numpy as np
from typing import Dict, List, Tuple, Optional


def calculate_accuracy(
    y_true: np.ndarray,
    y_pred: np.ndarray,
) -> float:
    """
    Рассчитать точность классификации (accuracy).

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).

    Returns:
        Доля правильных предсказаний (0–1).
    """
    return float(np.mean(y_true == y_pred))


def calculate_confusion_matrix(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    num_classes: int = 3,
) -> np.ndarray:
    """
    Рассчитать матрицу ошибок (confusion matrix).

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).
        num_classes: количество классов.

    Returns:
        Матрица ошибок формы (num_classes, num_classes).
    """
    cm = np.zeros((num_classes, num_classes), dtype=np.int64)
    for t, p in zip(y_true, y_pred):
        cm[t, p] += 1
    return cm


def calculate_precision_recall_f1(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    num_classes: int = 3,
) -> Dict[int, Dict[str, float]]:
    """
    Рассчитать precision, recall и F1-меру для каждого класса.

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).
        num_classes: количество классов.

    Returns:
        Словарь {класс: {'precision': ..., 'recall': ..., 'f1': ...}}.
    """
    cm = calculate_confusion_matrix(y_true, y_pred, num_classes)
    result: Dict[int, Dict[str, float]] = {}

    for c in range(num_classes):
        tp = cm[c, c]
        fp = cm[:, c].sum() - tp
        fn = cm[c, :].sum() - tp

        precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0
        recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0
        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0.0

        result[c] = {
            'precision': round(precision, 4),
            'recall': round(recall, 4),
            'f1': round(f1, 4),
        }

    return result


def calculate_macro_f1(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    num_classes: int = 3,
) -> float:
    """
    Рассчитать macro F1 (среднее F1 по классам).

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).
        num_classes: количество классов.

    Returns:
        Macro F1 score.
    """
    metrics = calculate_precision_recall_f1(y_true, y_pred, num_classes)
    return float(np.mean([m['f1'] for m in metrics.values()]))


def calculate_weighted_f1(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    num_classes: int = 3,
) -> float:
    """
    Рассчитать weighted F1 (взвешенный по количеству семплов в классе).

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).
        num_classes: количество классов.

    Returns:
        Weighted F1 score.
    """
    metrics = calculate_precision_recall_f1(y_true, y_pred, num_classes)
    class_counts = np.bincount(y_true, minlength=num_classes)
    total = class_counts.sum()

    weighted_f1 = sum(
        metrics[c]['f1'] * class_counts[c] / total
        for c in range(num_classes) if total > 0
    )
    return float(weighted_f1)


def classification_report(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    class_names: Optional[Dict[int, str]] = None,
) -> str:
    """
    Сформировать текстовый отчёт по метрикам классификации.

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).
        class_names: словарь {класс: имя} для отображения.

    Returns:
        Отформатированная строка с отчётом.
    """
    if class_names is None:
        class_names = {0: 'HOLD', 1: 'BUY', 2: 'SELL'}

    num_classes = len(class_names)
    cm = calculate_confusion_matrix(y_true, y_pred, num_classes)
    metrics = calculate_precision_recall_f1(y_true, y_pred, num_classes)

    accuracy = calculate_accuracy(y_true, y_pred)
    macro_f1 = calculate_macro_f1(y_true, y_pred, num_classes)
    weighted_f1 = calculate_weighted_f1(y_true, y_pred, num_classes)

    lines = []
    lines.append("=" * 70)
    lines.append("ОТЧЁТ ПО МЕТРИКАМ КЛАССИФИКАЦИИ")
    lines.append("=" * 70)
    lines.append(f"{'Класс':<10} {'Точность':>10} {'Полнота':>10} {'F1':>10} {'Поддержка':>12}")
    lines.append("-" * 70)

    for c in range(num_classes):
        name = class_names.get(c, str(c))
        count = int((y_true == c).sum())
        m = metrics[c]
        lines.append(
            f"{name:<10} {m['precision']:>10.4f} {m['recall']:>10.4f} {m['f1']:>10.4f} {count:>12}"
        )

    lines.append("-" * 70)
    lines.append(f"{'Accuracy':<10} {accuracy:>10.4f}")
    lines.append(f"{'Macro F1':<10} {macro_f1:>10.4f}")
    lines.append(f"{'Weighted F1':<10} {weighted_f1:>10.4f}")
    lines.append("=" * 70)

    lines.append("\nMATРИЦА ОШИБОК (Confusion Matrix):")
    lines.append(f"{'':>10}", end="")
    for c in range(num_classes):
        name = class_names.get(c, str(c))
        lines[-1] += f"{name:>8}"
    lines.append("")
    for i in range(num_classes):
        name = class_names.get(i, str(i))
        row = f"{name:<10}"
        for j in range(num_classes):
            row += f"{cm[i, j]:>8}"
        lines.append(row)

    return "\n".join(lines)


def calculate_all_metrics(
    y_true: np.ndarray,
    y_pred: np.ndarray,
    num_classes: int = 3,
) -> Dict[str, object]:
    """
    Рассчитать все метрики классификации.

    Args:
        y_true: истинные метки (N,).
        y_pred: предсказанные метки (N,).
        num_classes: количество классов.

    Returns:
        Словарь со всеми метриками.
    """
    return {
        'accuracy': calculate_accuracy(y_true, y_pred),
        'confusion_matrix': calculate_confusion_matrix(y_true, y_pred, num_classes).tolist(),
        'per_class': calculate_precision_recall_f1(y_true, y_pred, num_classes),
        'macro_f1': calculate_macro_f1(y_true, y_pred, num_classes),
        'weighted_f1': calculate_weighted_f1(y_true, y_pred, num_classes),
    }
