"""Validation metrics and ranking for normalized-sigmoid choices."""

import math
from collections.abc import Sequence
from dataclasses import dataclass
from enum import StrEnum
from itertools import groupby
from operator import itemgetter
from typing import Final, assert_never

import numpy as np
import torch

from raina_laya.features.grade_model.application.evaluate import evaluate_logits
from raina_laya.features.grade_model.application.evaluation_types import (
    CheckpointSelectionError,
    EvaluationInputError,
    EvaluationRole,
    EvaluationRoleError,
    ValidationMetrics,
)
from raina_laya.features.grade_model.domain.hierarchical import (
    ChoiceModelFamily,
    hierarchical_log_probs,
    hierarchical_predictions,
)
from raina_laya.features.grade_model.domain.v4_losses import choice_log_probs


@dataclass(frozen=True, slots=True)
class ChoiceCheckpoint:
    """Epoch and three-way validation metrics for a saved choice head."""

    epoch: int
    metrics: ValidationMetrics

    def __post_init__(self) -> None:
        """Keep every selection and frontier path validation-only."""
        if self.metrics.role is not EvaluationRole.VALIDATION:
            raise EvaluationRoleError(role=self.metrics.role)


class ChoiceSelectionPriority(StrEnum):
    """Choose the historical or constrained validation ranking."""

    GRADE = "grade"
    BINARY = "pass_fail"
    FAIL_MISS = "fail_miss"
    RECALL_FIRST = "pass_recall"


@dataclass(frozen=True, slots=True)
class PassFailMetrics:
    """Binary metrics projected from the unchanged three-grade predictions."""

    true_positives: int
    false_negatives: int
    false_positives: int
    true_negatives: int
    binary_macro_f1: float
    fail_recall: float
    pass_recall: float
    false_positive_rate: float


def evaluate_choices(
    logits: torch.Tensor,
    targets: torch.Tensor,
    role: EvaluationRole = EvaluationRole.VALIDATION,
    *,
    model_family: ChoiceModelFamily = ChoiceModelFamily.LEGACY,
) -> ValidationMetrics:
    """Evaluate family probabilities separately from its hard decision rule."""
    match model_family:
        case ChoiceModelFamily.LEGACY:
            return evaluate_logits(
                choice_log_probs(logits), targets, role, prediction_logits=logits
            )
        case ChoiceModelFamily.HIERARCHICAL:
            # One-hot decision indicators use the existing argmax hook honestly:
            # these are not raw scores or joint probabilities.
            decisions = torch.nn.functional.one_hot(
                hierarchical_predictions(logits),
                num_classes=3,
            ).to(logits.dtype)
            return evaluate_logits(
                hierarchical_log_probs(logits),
                targets,
                role,
                prediction_logits=decisions,
            )
        case unreachable:
            assert_never(unreachable)


def pass_fail_metrics(metrics: ValidationMetrics) -> PassFailMetrics:
    """Collapse A/S predictions into Pass without changing the deployed argmax."""
    fail, a, s = metrics.confusion_matrix
    tp = fail[0]
    fn = fail[1] + fail[2]
    fp = a[0] + s[0]
    tn = a[1] + a[2] + s[1] + s[2]
    if tp + fn == 0 or tn + fp == 0:
        raise CheckpointSelectionError(
            detail="validation needs both Fail and Pass for binary ranking"
        )
    return PassFailMetrics(
        tp,
        fn,
        fp,
        tn,
        (2 * tp / (2 * tp + fp + fn) + 2 * tn / (2 * tn + fp + fn)) / 2,
        tp / (tp + fn),
        tn / (tn + fp),
        fp / (tn + fp),
    )


GATE_CURVE_LEVELS: Final = (0.0, 0.005, 0.01, 0.02, 0.05)


@dataclass(frozen=True, slots=True)
class GateCurvePoint:
    """Gate errors at the threshold that caps Fail->Pass at one level.

    A same-set descriptive point, not an operating setting.
    """

    level: float
    threshold: float
    fail_to_pass_count: int
    fail_count: int
    pass_to_fail_count: int
    pass_count: int


def gate_auc(margins: Sequence[float], targets: Sequence[int]) -> float:
    """Return the ROC AUC of the gate margin, actual Pass (A or S) positive.

    A tied Pass/Fail pair counts half. Pairs are counted per tie group in exact
    integer arithmetic after one sort.
    """
    rows = sorted(
        (margin, target != 0) for margin, target in zip(margins, targets, strict=True)
    )
    passes = sum(is_pass for _margin, is_pass in rows)
    fails = len(rows) - passes
    if passes == 0 or fails == 0:
        raise EvaluationInputError(detail="gate AUC needs both Fail and Pass songs")
    fails_below = twice_wins = 0
    for _margin, group in groupby(rows, key=itemgetter(0)):
        tied = [is_pass for _score, is_pass in group]
        tied_passes = sum(tied)
        tied_fails = len(tied) - tied_passes
        twice_wins += tied_passes * (2 * fails_below + tied_fails)
        fails_below += tied_fails
    return twice_wins / (2 * passes * fails)


def gate_curve(
    margins: Sequence[float],
    targets: Sequence[int],
    levels: Sequence[float] = GATE_CURVE_LEVELS,
) -> list[GateCurvePoint]:
    """Read gate error pairs off this same set's own Fail margins.

    The threshold of a level is the (floor(level * fail_count) + 1)-th highest Fail
    margin and a song is Pass iff its margin is strictly above it. Fail->Pass is
    therefore at most floor(level * fail_count) (fewer when Fail margins tie at the
    threshold) and exactly 0 at level 0. These are same-set descriptive curve points,
    not operating settings: each threshold is fitted to the songs it scores.
    """
    fail_margins = sorted(
        (m for m, target in zip(margins, targets, strict=True) if target == 0),
        reverse=True,
    )
    pass_margins = [m for m, target in zip(margins, targets, strict=True) if target]
    if not fail_margins or not pass_margins:
        raise EvaluationInputError(detail="gate curve needs both Fail and Pass songs")
    points: list[GateCurvePoint] = []
    for level in levels:
        threshold = fail_margins[math.floor(level * len(fail_margins))]
        points.append(
            GateCurvePoint(
                level,
                threshold,
                sum(m > threshold for m in fail_margins),
                len(fail_margins),
                sum(m <= threshold for m in pass_margins),
                len(pass_margins),
            )
        )
    return points


@dataclass(frozen=True, slots=True)
class RejectBand:
    """Asymmetric gate-margin band: Pass iff margin >= hi, Fail iff margin <= lo."""

    lo: float
    hi: float


def reject_band_counts(
    margins: Sequence[float], targets: Sequence[int], band: RejectBand
) -> dict[str, float]:
    """Decided-song error counts once songs with lo < margin < hi are rejected."""
    m = np.asarray(margins, dtype=np.float64)
    is_pass = np.asarray(targets) != 0
    rejected = (m > band.lo) & (m < band.hi)
    f2p = int((~is_pass & (m >= band.hi)).sum())
    fail_count = int((~is_pass & ~rejected).sum())
    p2f = int((is_pass & (m <= band.lo)).sum())
    pass_count = int((is_pass & ~rejected).sum())
    rejected_pass = int((is_pass & rejected).sum())
    return {
        "fail_to_pass_count": f2p,
        "fail_count": fail_count,
        "fail_to_pass_rate": f2p / fail_count if fail_count else 0.0,
        "pass_to_fail_count": p2f,
        "pass_correct_count": pass_count - p2f,
        "pass_count": pass_count,
        "pass_to_fail_rate": p2f / pass_count if pass_count else 0.0,
        "rejected_total": int(rejected.sum()),
        "rejected_pass_count": rejected_pass,
        "rejected_pass_rate": rejected_pass / max(int(is_pass.sum()), 1),
    }


# Share of actual Pass songs a band may reject (user decisions on 2026-10-07: 30%,
# then 50%, then 70% at 23:1x KST: "no Fail->Pass at all comes first"; see
# _Autoresearch/20261007/06_거부_구간_계획.md §7 and AGENTS.md 1-1).
MAX_PASS_LOSS: Final = 0.70


def select_reject_band(
    margins: Sequence[float],
    targets: Sequence[int],
    max_pass_loss: float = MAX_PASS_LOSS,
    grid: int = 200,
) -> RejectBand:
    """Choose (lo, hi) on one set's own margins; call it on validation only.

    Feasible bands reject at most `max_pass_loss` of the actual Pass songs and keep
    the decided Pass->Fail rate at or below the no-reject rate at boundary 0 (so a
    band never buys Fail->Pass by shifting errors onto Pass->Fail; without this cap
    hi = +inf is always "optimal"). Among feasible bands: fewest decided Fail->Pass,
    then fewest decided Pass->Fail, then fewest rejected songs. Cut points are the
    margin quantiles plus 0. Falls back to (0, 0) when nothing is feasible.
    """
    m = np.asarray(margins, dtype=np.float64)
    is_pass = np.asarray(targets) != 0
    if m.ndim != 1 or m.shape != is_pass.shape or m.size == 0:
        raise EvaluationInputError(detail="reject band needs equal-length 1-D inputs")
    pm, fm = np.sort(m[is_pass]), np.sort(m[~is_pass])
    if pm.size == 0 or fm.size == 0:
        raise EvaluationInputError(detail="reject band needs both Fail and Pass songs")
    max_p2f_rate = (pm <= 0.0).sum() / pm.size
    cuts = np.unique(np.concatenate([np.quantile(m, np.linspace(0, 1, grid)), [0.0]]))
    lo, hi = cuts[:, None], cuts[None, :]
    pass_le_lo = np.searchsorted(pm, lo, side="right")  # Pass decided Fail
    pass_ge_hi = pm.size - np.searchsorted(pm, hi, side="left")  # Pass decided Pass
    fail_ge_hi = fm.size - np.searchsorted(fm, hi, side="left")  # Fail decided Pass
    fail_le_lo = np.searchsorted(fm, lo, side="right")
    decided_pass = pass_le_lo + pass_ge_hi
    rejected = (pm.size - decided_pass) + (fm.size - fail_le_lo - fail_ge_hi)
    p2f_rate = np.divide(
        pass_le_lo,
        decided_pass,
        out=np.ones(decided_pass.shape),
        where=decided_pass > 0,
    )
    feasible = (
        (lo <= hi)
        & (pm.size - decided_pass <= max_pass_loss * pm.size + 1e-9)
        & (p2f_rate <= max_p2f_rate + 1e-12)
    )
    if not feasible.any():
        return RejectBand(lo=0.0, hi=0.0)
    big = 1 << 40
    keys = [
        np.where(feasible, k, big).ravel() for k in (fail_ge_hi, pass_le_lo, rejected)
    ]
    i, j = np.unravel_index(np.lexsort(keys[::-1])[0], feasible.shape)
    return RejectBand(lo=float(cuts[i]), hi=float(cuts[j]))


def update_choice_frontier(
    frontier: tuple[ChoiceCheckpoint, ...],
    candidate: ChoiceCheckpoint,
    minimum_pass_recall: float,
) -> tuple[ChoiceCheckpoint, ...]:
    """Retain eligible nondominated false-pass and Pass-recall coordinates."""
    proposed = pass_fail_metrics(candidate.metrics)
    if proposed.pass_recall < minimum_pass_recall:
        return frontier
    existing = tuple(pass_fail_metrics(point.metrics) for point in frontier)
    if any(
        point.false_negatives <= proposed.false_negatives
        and point.true_negatives >= proposed.true_negatives
        for point in existing
    ):
        return frontier
    retained = tuple(
        checkpoint
        for checkpoint, point in zip(frontier, existing, strict=True)
        if not (
            proposed.false_negatives <= point.false_negatives
            and proposed.true_negatives >= point.true_negatives
        )
    )
    return (*retained, candidate)


def choice_checkpoint_key(
    checkpoint: ChoiceCheckpoint,
    priority: ChoiceSelectionPriority = ChoiceSelectionPriority.GRADE,
    minimum_pass_recall: float = 0.79,
    *,
    maximum_fail_to_pass_rate: float = 0.05,
) -> tuple[float, ...]:
    """Rank validated checkpoints under the explicitly selected research goal."""
    metrics = checkpoint.metrics
    macro_f1 = metrics.macro_f1
    s_f1 = metrics.per_class[2].f1
    nll = metrics.nll
    if macro_f1 is None or s_f1 is None or nll is None:
        raise CheckpointSelectionError(
            detail="validation needs all three grades to rank checkpoints"
        )
    match priority:
        case ChoiceSelectionPriority.GRADE:
            return macro_f1, s_f1, -nll
        case ChoiceSelectionPriority.BINARY:
            binary = pass_fail_metrics(metrics)
            return (
                binary.binary_macro_f1,
                binary.fail_recall,
                binary.pass_recall,
                macro_f1,
                s_f1,
                -nll,
            )
        case ChoiceSelectionPriority.FAIL_MISS:
            binary = pass_fail_metrics(metrics)
            if binary.pass_recall < minimum_pass_recall:
                raise CheckpointSelectionError(
                    detail="validation Pass recall is below the required floor"
                )
            return binary.fail_recall, binary.binary_macro_f1, macro_f1, s_f1, -nll
        case ChoiceSelectionPriority.RECALL_FIRST:
            binary = pass_fail_metrics(metrics)
            fail_to_pass_rate = binary.false_negatives / (
                binary.true_positives + binary.false_negatives
            )
            if (
                binary.pass_recall < minimum_pass_recall
                or fail_to_pass_rate >= maximum_fail_to_pass_rate
            ):
                raise CheckpointSelectionError(
                    detail="validation violates recall-first selection limits"
                )
            return binary.pass_recall, binary.fail_recall, macro_f1, s_f1, -nll
        case unreachable:
            assert_never(unreachable)
