"""Deterministic v1.4 branch correlation for weekly wizard evidence."""

from __future__ import annotations

from collections.abc import Callable, Mapping
from typing import TypeVar

from .models import Event, TrainerSessionPayload
from .weekly_operations_wizard_host_events_r4 import (
    ProjectionSource,
    WIZARD_MAPPING_ADAPTER,
    WizardValue,
    event_value,
    numeric_float,
    project_customer_event_state,
    string_mapping,
)
from .weekly_operations_wizard_host_types_r4 import BranchDecision, SafetyClassification
from .wizard_models import WizardBranch, WizardFlow

SafetyClassifier = Callable[[WizardFlow, Mapping[str, WizardValue]], SafetyClassification]
TrainerEvidence = TrainerSessionPayload | Event | Mapping[str, WizardValue]
AnswerT = TypeVar("AnswerT")
TrainerT = TypeVar("TrainerT")


def _float(value: WizardValue) -> float | None:
    return numeric_float(value)


def _int(value: WizardValue) -> int | None:
    if not isinstance(value, (bool, int, float, str)):
        return None
    try:
        return int(value)
    except ValueError:
        return None


def _trainer_value(source: TrainerEvidence | None, *names: str) -> WizardValue:
    if source is None:
        return None
    return event_value(source, *names)


def _trainer_evidence(value: WizardValue) -> TrainerEvidence | None:
    if isinstance(value, (TrainerSessionPayload, Event)) or string_mapping(value):
        return value
    return None


def deterministic_weekly_branch(
    answers: Mapping[str, AnswerT],
    *,
    classifier: SafetyClassifier,
    prior_week_same_weekday_weight: float | None,
    trainer_session: TrainerSessionPayload | Event | Mapping[str, TrainerT] | None,
    safety: SafetyClassification | None,
    canonical_event_state: ProjectionSource | None,
    aliases: Mapping[str, WizardValue],
) -> BranchDecision:
    """Apply the fixed sleep, condition, performance, and change branch table."""
    parsed_answers = WIZARD_MAPPING_ADAPTER.validate_python(answers)
    parsed_trainer: TrainerEvidence | None
    match trainer_session:
        case TrainerSessionPayload() | Event():
            parsed_trainer = trainer_session
        case Mapping():
            parsed_trainer = WIZARD_MAPPING_ADAPTER.validate_python(trainer_session)
        case None:
            parsed_trainer = None
    if canonical_event_state is not None:
        state_day = aliases.get("kst_day") or parsed_answers.get("kst_day")
        projected = project_customer_event_state(
            canonical_event_state,
            str(state_day) if state_day is not None else None,
        )
        if prior_week_same_weekday_weight is None:
            prior_week_same_weekday_weight = _float(
                projected.get("prior_week_same_weekday_weight")
            )
        if parsed_trainer is None:
            parsed_trainer = _trainer_evidence(projected.get("trainer_session"))
    current = _float(
        aliases.get("bodyweight_kg")
        or parsed_answers.get("bodyweight")
        or parsed_answers.get("body_weight_kg")
    )
    if prior_week_same_weekday_weight is None:
        for name in (
            "prior_weight_kg",
            "previous_weight_kg",
            "prior_week_weight",
            "previous_week_weight",
        ):
            prior_week_same_weekday_weight = _float(aliases.get(name))
            if prior_week_same_weekday_weight is not None:
                break
    followups: list[str] = []
    reasons: list[str] = []
    sleep_hours = _float(
        parsed_answers.get("sleep_duration", parsed_answers.get("sleep_hours"))
    )
    sleep_quality = _int(
        parsed_answers.get("sleep_quality", parsed_answers.get("sleep_quality_1to5"))
    )
    condition = _int(
        parsed_answers.get("condition", parsed_answers.get("readiness_1to5"))
    )
    completion = str(
        parsed_answers.get(
            "completion",
            parsed_answers.get("workout_completion", parsed_answers.get("done", "")),
        )
    ).lower()
    performance = _int(
        parsed_answers.get("workout_quality")
        or parsed_answers.get("workout_quality_1to5")
        or _trainer_value(parsed_trainer, "performance", "performance_1to5")
        or parsed_answers.get("performance")
    )
    trainer_intensity = str(
        _trainer_value(parsed_trainer, "intensity", "intensity_vs_plan")
        or parsed_answers.get("intensity", "")
    ).lower()
    trainer_pain = str(
        _trainer_value(parsed_trainer, "pain", "pain_summary")
        or parsed_answers.get("pain", "")
    ).lower()
    trainer_below = trainer_intensity == "below"
    if (sleep_hours is not None and sleep_hours < 6.0) or (
        sleep_quality is not None and sleep_quality <= 2
    ):
        followups.extend(("Q-SLEEP-CAUSE", "Q-SLEEP-ADJUST"))
        reasons.append("sleep_anomaly")
    if condition is not None and condition <= 2:
        followups.extend(("Q-COND-SYMPTOM", "Q-COND-INTENSITY"))
        reasons.append("condition_anomaly")
    if (
        completion
        in {"missed", "partial", "rest_changed", "rest", "not_done", "false", "no"}
        or (performance is not None and performance <= 2)
        or trainer_below
    ):
        followups.extend(("Q-PERF-REASON", "Q-PERF-NEXT"))
        reasons.append("performance_anomaly")
    change = (
        current is not None
        and prior_week_same_weekday_weight is not None
        and abs(current - prior_week_same_weekday_weight) >= 1.0
    )
    trainer_change = bool(
        trainer_pain and trainer_pain not in {"none", "normal", "없음"}
    ) or trainer_below
    if change:
        reasons.append("weight_change")
    if trainer_change:
        reasons.append("trainer_change")
    trainer_flow = any(
        key in parsed_answers
        for key in ("done", "performance", "intensity", "operator_note")
    )
    classification = safety or classifier(
        WizardFlow.TRAINER_SESSION if trainer_flow else WizardFlow.NUTRITION,
        parsed_answers,
    )
    unique_reasons = tuple(dict.fromkeys(reasons))
    unique_followups = tuple(dict.fromkeys(followups))
    if classification:
        return BranchDecision(
            WizardBranch.SAFETY_HOLD,
            (),
            True,
            (*unique_reasons, "safety_hold"),
        )
    if change or trainer_change:
        return BranchDecision(WizardBranch.CHANGE, unique_followups, True, unique_reasons)
    if unique_followups:
        return BranchDecision(WizardBranch.ANOMALY, unique_followups, True, unique_reasons)
    return BranchDecision(WizardBranch.NORMAL)
