"""Mean-gate-margin ensemble scoring of several v5 choice checkpoints on one role."""

import hashlib
import json
from collections.abc import Sequence
from contextlib import nullcontext
from dataclasses import dataclass, replace
from pathlib import Path
from typing import cast

import torch

from raina_laya.features.grade_model.client import (
    ChoiceHeadCheckpoint,
    ChoiceModelFamily,
    ChoiceScores,
    ChoiceTrainConfig,
    EvaluationRole,
    combine_choice_scores,
    evaluate_choices,
    final_train_statistics,
    hierarchical_gate_margin,
    load_choice_head,
    load_choice_scores,
    publish_choice_scores,
    validate_choice_config,
)
from raina_laya.workflows.train_grade_model import TrainingRunInputError
from raina_laya.workflows.v4_choice import _healthy_provider
from raina_laya.workflows.v4_choice_banks import (
    EvaluationBankInputs,
    EvaluationBankPool,
)
from raina_laya.workflows.v4_choice_heldout import (
    HeldOutPaths,
    _checkpoint_payload,
    _clear_failure_frames,
    _EvaluationInputs,
    _require_lineage,
    _write_scores,
    gate_report,
    metrics_report,
    reject_report,
)
from raina_laya.workflows.v4_choice_metadata import choice_model_metadata
from raina_laya.workflows.v4_choice_score_cache import score_cache_context
from raina_laya.workflows.v4_choice_validation import score_choices

__all__ = [
    "EnsembleEvaluationJob",
    "EnsembleMember",
    "ensemble_checkpoint_sha256",
    "ensemble_logits",
    "load_manifest",
    "run_ensemble_evaluation",
    "run_ensemble_evaluations",
]


@dataclass(frozen=True, slots=True)
class EnsembleMember:
    """One member's parsed training config and its checkpoint, with config paths."""

    config_path: Path
    config: ChoiceTrainConfig
    paths: HeldOutPaths


@dataclass(frozen=True, slots=True)
class EnsembleEvaluationJob:
    """One ordered ensemble and its exclusive report and optional score outputs."""

    members: list[EnsembleMember]
    manifest: Path
    output: Path
    scores_output: Path | None = None
    band_scores: Path | None = None


@dataclass(frozen=True, slots=True)
class _EnsembleInputs:
    owned: _EvaluationInputs
    banks: EvaluationBankPool


def load_manifest(path: Path) -> list[tuple[Path, Path]]:
    """Return the (config, checkpoint) pairs of a manifest with at least two members."""
    manifest = json.loads(path.read_text(encoding="utf-8"))
    members = manifest.get("members") if isinstance(manifest, dict) else None
    if (
        not isinstance(members, list)
        or len(members) < 2  # noqa: PLR2004 - an ensemble needs two or more members
        or not all(
            isinstance(member, dict)
            and isinstance(member.get("config"), str)
            and isinstance(member.get("checkpoint"), str)
            for member in members
        )
    ):
        raise TrainingRunInputError(
            detail='manifest needs "members": two or more {"config","checkpoint"}'
        )
    return [(Path(member["config"]), Path(member["checkpoint"])) for member in members]


def ensemble_checkpoint_sha256(member_digests: list[str]) -> str:
    """Digest of the member checkpoint digests, concatenated in manifest order."""
    return hashlib.sha256("".join(member_digests).encode()).hexdigest()


def ensemble_logits(
    margins: list[torch.Tensor], conditionals: list[torch.Tensor]
) -> torch.Tensor:
    """Per-song mean gate margin m and conditional logit c as [-m/2, +m/2, c]."""
    return combine_choice_scores(margins, conditionals)


def _check_shared_inputs(members: list[EnsembleMember]) -> None:
    first = members[0].paths
    for member in members:
        if member.config.model_family is not ChoiceModelFamily.HIERARCHICAL:
            raise TrainingRunInputError(detail="ensemble needs the hierarchical family")
        validate_choice_config(member.config)
        if (member.paths.split, member.paths.data_root) != (
            first.split,
            first.data_root,
        ):
            raise TrainingRunInputError(
                detail="ensemble members must share the split and data_root"
            )


@dataclass(frozen=True, slots=True)
class _Scored:
    margin: torch.Tensor
    conditional: torch.Tensor
    targets: torch.Tensor
    source_identifiers: tuple[str, ...]
    split_digest: str
    cache_inventory_digest: str
    sha256: str
    epoch: int


def _score_member(  # noqa: PLR0913, PLR0917 - shared invocation contexts and optional cache
    member: EnsembleMember,
    role: EvaluationRole,
    device: torch.device,
    owned: _EvaluationInputs,
    banks: EvaluationBankPool,
    score_cache_dir: Path | None,
) -> _Scored:
    """Score one member on the role songs of its own feature setting."""
    paths, config = member.paths, member.config
    raw = paths.checkpoint.read_bytes()
    payload = _checkpoint_payload(raw)
    layers = config.embedding_layers or ()
    split_file_sha256 = hashlib.sha256(paths.split.read_bytes()).hexdigest()
    dataset = owned.load(config, paths, role, banks, _healthy_provider)
    statistics = final_train_statistics(
        dataset.development_targets,
        config.micro_batch,
        config.rare_weight_gamma,
        config.class_weight_grouping,
    )
    epoch = _require_lineage(payload, config, dataset, statistics.counts)
    if score_cache_dir is not None and owned.score_context is None:
        owned.score_context = score_cache_context(
            config,
            paths.split,
            dataset.split_digest,
            dataset.cache_inventory_digest,
            dataset.feature_policy,
            role,
            dataset.test_songs,
            dataset.test_ids,
            device=device,
            feature_content=dataset.feature_content,
        )
    context = owned.score_context
    if context is not None:
        context.require_matching_inputs(dataset.test_songs)
    # Configs may share feature inputs while differing in architecture metadata.
    identity = (
        replace(
            context.identity(hashlib.sha256(raw).hexdigest()),
            architecture_json=json.dumps(choice_model_metadata(config), sort_keys=True),
        )
        if context is not None
        else None
    )
    cached = (
        load_choice_scores(score_cache_dir, identity)
        if score_cache_dir is not None and identity is not None
        else None
    )
    inputs = EvaluationBankInputs(
        dataset,
        split_file_sha256,
        role,
        paths.layer_dir or paths.database or paths.inventory,
        layers,
    )
    if cached is not None:
        logits, targets = cached.logits, cached.targets
    else:
        model = load_choice_head(
            config, statistics, cast("ChoiceHeadCheckpoint", payload)
        ).to(device)
        with banks.acquire(inputs) as bank:
            logits, targets = (
                score_choices(model, dataset.test_songs, device, bank=bank)
                if bank is not None
                else score_choices(model, dataset.test_songs, device)
            )
        if score_cache_dir is not None and identity is not None:
            _ = publish_choice_scores(
                score_cache_dir, ChoiceScores(identity, logits, targets)
            )
    return _Scored(
        hierarchical_gate_margin(logits),
        logits[:, 2],
        targets,
        dataset.test_ids,
        dataset.split_digest,
        dataset.cache_inventory_digest,
        hashlib.sha256(raw).hexdigest(),
        epoch,
    )


def run_ensemble_evaluation(  # noqa: PLR0913, PLR0917 - mirrors the held-out inputs
    members: list[EnsembleMember],
    manifest: Path,
    output: Path,
    device: torch.device,
    role: EvaluationRole = EvaluationRole.TEST,
    scores_output: Path | None = None,
    band_scores: Path | None = None,
    reuse_banks: bool | None = None,
    bank_memory_budget_bytes: int | None = None,
    score_cache_dir: Path | None = None,
) -> dict[str, object]:
    """Score the mean-gate-margin ensemble on `role` songs and publish the report.

    The ensemble predicts Pass iff the mean member gate margin is above zero, then S
    iff the mean conditional logit is above zero; the report has the keys of a single
    checkpoint evaluation. Its `checkpoint_sha256` is the sha256 of the member
    checkpoint sha256 hex strings concatenated in manifest order and `epoch` is null.
    Outputs are exclusive; the per-song mean margins are written before the report.
    Raised errors retain diagnostic stack entries while completed frames release
    their owned input/model/bank references before the error leaves the workflow.
    """
    try:
        return _run_ensemble_evaluation(
            members,
            manifest,
            output,
            device,
            role,
            scores_output,
            band_scores,
            reuse_banks,
            bank_memory_budget_bytes,
            score_cache_dir,
        )
    except BaseException as error:
        _clear_failure_frames(error)
        raise


def _run_ensemble_evaluation(  # noqa: PLR0913, PLR0917 - preserves public inputs
    members: list[EnsembleMember],
    manifest: Path,
    output: Path,
    device: torch.device,
    role: EvaluationRole,
    scores_output: Path | None,
    band_scores: Path | None,
    reuse_banks: bool | None,
    bank_memory_budget_bytes: int | None,
    score_cache_dir: Path | None,
    shared: _EnsembleInputs | None = None,
) -> dict[str, object]:
    """Keep owned inputs in a frame that completes before raised-error cleanup."""
    if len(members) < 2:  # noqa: PLR2004 - an ensemble needs two or more members
        raise TrainingRunInputError(detail="ensemble needs two or more members")
    _check_shared_inputs(members)
    for existing in (output, scores_output):
        if existing is not None and existing.exists():
            raise FileExistsError(existing)
    owned = shared.owned if shared is not None else _EvaluationInputs()
    pool = (
        nullcontext(shared.banks)
        if shared is not None
        else EvaluationBankPool(
            device, enabled=reuse_banks, budget_bytes=bank_memory_budget_bytes
        )
    )
    with pool as banks:
        scored = [
            _score_member(member, role, device, owned, banks, score_cache_dir)
            for member in members
        ]
    first, targets = scored[0], scored[0].targets
    for other in scored[1:]:
        if (
            other.split_digest != first.split_digest
            or other.source_identifiers != first.source_identifiers
            or not torch.equal(other.targets, targets)
        ):
            raise TrainingRunInputError(
                detail="ensemble members must score the same songs in the same order"
            )
    digests = [one.sha256 for one in scored]
    logits = ensemble_logits(
        [one.margin for one in scored], [one.conditional for one in scored]
    )
    metrics = evaluate_choices(
        logits, targets, role, model_family=ChoiceModelFamily.HIERARCHICAL
    )
    mean_margins = hierarchical_gate_margin(logits).tolist()
    labels = targets.tolist()
    split = members[0].paths.split
    report: dict[str, object] = {
        "ensemble": "mean_gate_margin",
        "manifest_sha256": hashlib.sha256(manifest.read_bytes()).hexdigest(),
        "members": [
            {
                "config": str(member.config_path),
                "checkpoint": str(member.paths.checkpoint),
                "checkpoint_sha256": one.sha256,
                "epoch": one.epoch,
                "model_family": member.config.model_family.value,
                "embedding_layers": (
                    None
                    if member.config.embedding_layers is None
                    else list(member.config.embedding_layers)
                ),
            }
            for member, one in zip(members, scored, strict=True)
        ],
        "checkpoint": str(manifest),  # key parity with a single-checkpoint report
        "checkpoint_sha256": ensemble_checkpoint_sha256(digests),
        "epoch": None,
        "model_family": ChoiceModelFamily.HIERARCHICAL.value,
        "split": str(split),
        "split_file_sha256": hashlib.sha256(split.read_bytes()).hexdigest(),
        "split_digest": first.split_digest,
        # The first member's inventory digest; members may use other inventories.
        "cache_inventory_digest": first.cache_inventory_digest,
        "test_songs": len(first.source_identifiers),
        **metrics_report(metrics),
        **gate_report(mean_margins, labels),
    }
    if band_scores is not None:
        report.update(reject_report(mean_margins, labels, band_scores))
    if scores_output is not None:
        _write_scores(scores_output, first.source_identifiers, labels, mean_margins)
    output.parent.mkdir(parents=True, exist_ok=True)
    with output.open("x", encoding="utf-8") as stream:
        _ = stream.write(json.dumps(report, sort_keys=True) + "\n")
    return report


def run_ensemble_evaluations(  # noqa: PLR0913 - optional invocation controls
    jobs: Sequence[EnsembleEvaluationJob],
    device: torch.device,
    role: EvaluationRole = EvaluationRole.TEST,
    *,
    reuse_banks: bool | None = None,
    bank_memory_budget_bytes: int | None = None,
    score_cache_dir: Path | None = None,
) -> tuple[
    list[dict[str, object] | None], list[tuple[EnsembleEvaluationJob, Exception]]
]:
    """Evaluate ordered jobs sharing only verified invocation-owned feature inputs.

    Each member keeps its actual config and checkpoint lineage. A failed job leaves
    a None report and does not stop later jobs. Retained exceptions preserve their
    diagnostic stacks while completed frames release owned tensors.
    """
    try:
        reports, failures = _run_ensemble_evaluations(
            jobs, device, role, reuse_banks, bank_memory_budget_bytes, score_cache_dir
        )
    except BaseException as error:
        _clear_failure_frames(error)
        raise
    for _job, error in failures:
        _clear_failure_frames(error)
    return reports, failures


def _run_ensemble_evaluations(  # noqa: PLR0913, PLR0917 - invocation inputs
    jobs: Sequence[EnsembleEvaluationJob],
    device: torch.device,
    role: EvaluationRole,
    reuse_banks: bool | None,
    bank_memory_budget_bytes: int | None,
    score_cache_dir: Path | None,
) -> tuple[
    list[dict[str, object] | None], list[tuple[EnsembleEvaluationJob, Exception]]
]:
    if not jobs:
        raise TrainingRunInputError(detail="ensemble jobs lists no manifests")
    for job in jobs:
        if len(job.members) < 2:  # noqa: PLR2004 - two members form an ensemble
            raise TrainingRunInputError(detail="ensemble needs two or more members")
        _check_shared_inputs(job.members)
        for path in (job.output, job.scores_output):
            if path is not None and path.exists():
                raise FileExistsError(path)
    reports: list[dict[str, object] | None] = []
    failures: list[tuple[EnsembleEvaluationJob, Exception]] = []
    with EvaluationBankPool(
        device, enabled=reuse_banks, budget_bytes=bank_memory_budget_bytes
    ) as banks:
        shared = _EnsembleInputs(_EvaluationInputs(), banks)
        for job in jobs:
            try:
                reports.append(
                    _run_ensemble_evaluation(
                        job.members,
                        job.manifest,
                        job.output,
                        device,
                        role,
                        job.scores_output,
                        job.band_scores,
                        reuse_banks,
                        bank_memory_budget_bytes,
                        score_cache_dir,
                        shared,
                    )
                )
            except Exception as error:  # noqa: BLE001 - isolated failures returned to caller
                reports.append(None)
                failures.append((job, error))
                _clear_failure_frames(error)
        shared.owned.clear(banks)
    return reports, failures
