"""Canonical adherence evidence for weekly operations."""

from __future__ import annotations

import hashlib
import json
from collections.abc import Mapping
from dataclasses import dataclass
from decimal import Decimal, InvalidOperation
from typing import ClassVar, Final, Literal, TypeAlias, TypeGuard

from pydantic import BaseModel, ConfigDict, Field, model_validator

ADHERENCE_SIGNAL_SCHEMA_VERSION: Final = "1.0"
_DEFAULT_TOLERANCE: Final = Decimal("10")
MetricValue: TypeAlias = str | int | float | Decimal | None
AdherenceValue: TypeAlias = MetricValue | Mapping[str, MetricValue]
AdherenceStatus: TypeAlias = Literal[
    "adequate", "inadequate", "missing", "contradictory"
]


class AdherenceSignalError(ValueError):
    """Canonical adherence input or invariant is invalid."""


class CanonicalAdherenceSignal(BaseModel):
    """Versioned immutable evidence derived from canonical target and actual values."""

    model_config: ClassVar[ConfigDict] = ConfigDict(frozen=True, extra="forbid")

    schema_version: Literal["1.0"] = ADHERENCE_SIGNAL_SCHEMA_VERSION
    source_event_ids: tuple[str, ...] = ()
    target_calories_kcal: int | None = Field(default=None, ge=0, le=30000)
    actual_calories_kcal: int | None = Field(default=None, ge=0, le=30000)
    target_carbohydrate_g: int | None = Field(default=None, ge=0, le=2000)
    actual_carbohydrate_g: int | None = Field(default=None, ge=0, le=2000)
    target_protein_g: int | None = Field(default=None, ge=0, le=1000)
    actual_protein_g: int | None = Field(default=None, ge=0, le=1000)
    target_fat_g: int | None = Field(default=None, ge=0, le=1000)
    actual_fat_g: int | None = Field(default=None, ge=0, le=1000)
    tolerance_percent: Decimal = Field(default=_DEFAULT_TOLERANCE, ge=0, le=100)
    complete: bool = False
    adherent: bool | None = None
    status: AdherenceStatus = "missing"
    reason: str | None = Field(default=None, max_length=160)

    @model_validator(mode="after")
    def validate_signal(self) -> CanonicalAdherenceSignal:
        complete_status = self.status in {"adequate", "inadequate"}
        if complete_status != self.complete:
            raise AdherenceSignalError("adherence status and completeness disagree")
        if complete_status != (self.adherent is not None):
            raise AdherenceSignalError("adherence result and completeness disagree")
        targets = (
            self.target_calories_kcal,
            self.target_carbohydrate_g,
            self.target_protein_g,
            self.target_fat_g,
        )
        actuals = (
            self.actual_calories_kcal,
            self.actual_carbohydrate_g,
            self.actual_protein_g,
            self.actual_fat_g,
        )
        if any(target is not None and target <= 0 for target in targets):
            raise AdherenceSignalError("adherence targets must be positive when supplied")
        if (
            any(
                (target is None) != (actual is None)
                for target, actual in zip(targets, actuals, strict=True)
            )
            and self.status != "contradictory"
        ):
            raise AdherenceSignalError("partial adherence pairs must be contradictory")
        return self

    @property
    def version(self) -> str:
        return self.schema_version

    @property
    def digest(self) -> str:
        encoded = json.dumps(
            self.model_dump(mode="json"),
            ensure_ascii=False,
            sort_keys=True,
            separators=(",", ":"),
        )
        return hashlib.sha256(encoded.encode()).hexdigest()


AdherenceSignal = CanonicalAdherenceSignal


@dataclass(frozen=True, slots=True)
class _SignalValues:
    source_event_ids: tuple[str, ...]
    target_calories: int | None
    actual_calories: int | None
    target_carbohydrate: int | None
    actual_carbohydrate: int | None
    target_protein: int | None
    actual_protein: int | None
    target_fat: int | None
    actual_fat: int | None
    tolerance: Decimal


def _adherence_int(value: AdherenceValue) -> int | None:
    if value is None:
        return None
    try:
        parsed = int(Decimal(str(value)))
    except (InvalidOperation, ValueError, OverflowError):
        return None
    return parsed if parsed >= 0 else None


def _is_metric_mapping(value: AdherenceValue) -> TypeGuard[Mapping[str, MetricValue]]:
    return isinstance(value, Mapping)


def _nested_values(
    source: Mapping[str, AdherenceValue] | None,
    primary: str,
    secondary: str,
) -> Mapping[str, AdherenceValue]:
    if source is None:
        return {}
    candidate = source.get(primary)
    if _is_metric_mapping(candidate):
        return candidate
    candidate = source.get(secondary)
    return candidate if _is_metric_mapping(candidate) else source


def _metric_value(
    source: Mapping[str, AdherenceValue], names: tuple[str, ...]
) -> int | None:
    for name in names:
        if name in source:
            return _adherence_int(source[name])
    return None


def _signal(
    values: _SignalValues, status: AdherenceStatus, reason: str | None
) -> CanonicalAdherenceSignal:
    complete = status in {"adequate", "inadequate"}
    adherent = status == "adequate" if complete else None
    return CanonicalAdherenceSignal(
        schema_version=ADHERENCE_SIGNAL_SCHEMA_VERSION,
        source_event_ids=values.source_event_ids,
        target_calories_kcal=values.target_calories,
        actual_calories_kcal=values.actual_calories,
        target_carbohydrate_g=values.target_carbohydrate,
        actual_carbohydrate_g=values.actual_carbohydrate,
        target_protein_g=values.target_protein,
        actual_protein_g=values.actual_protein,
        target_fat_g=values.target_fat,
        actual_fat_g=values.actual_fat,
        tolerance_percent=values.tolerance,
        complete=complete,
        adherent=adherent,
        status=status,
        reason=reason,
    )


def derive_canonical_adherence_signal(
    target: Mapping[str, AdherenceValue] | None,
    actual: Mapping[str, AdherenceValue] | None,
    *,
    tolerance_percent: Decimal | int | str = _DEFAULT_TOLERANCE,
    source_event_ids: tuple[str, ...] = (),
) -> CanonicalAdherenceSignal:
    """Derive immutable evidence from canonical target and actual values."""
    try:
        tolerance = Decimal(str(tolerance_percent))
    except (InvalidOperation, ValueError) as error:
        raise AdherenceSignalError("adherence tolerance is invalid") from error
    if tolerance < 0 or tolerance > 100:
        raise AdherenceSignalError("adherence tolerance is invalid")
    target_map = _nested_values(target, "target", "check_in_target")
    actual_map = _nested_values(actual, "actual", "check_in")
    aliases: Final = {
        "calories": (
            "calories_kcal", "calories", "target_calories_kcal", "actual_calories_kcal"
        ),
        "carbohydrate": (
            "carbohydrate_g", "carbs_g", "carbs",
            "target_carbohydrate_g", "actual_carbohydrate_g",
        ),
        "protein": ("protein_g", "protein", "target_protein_g", "actual_protein_g"),
        "fat": ("fat_g", "fat", "target_fat_g", "actual_fat_g"),
    }
    targets = {name: _metric_value(target_map, names) for name, names in aliases.items()}
    actuals = {name: _metric_value(actual_map, names) for name, names in aliases.items()}
    pairs = tuple(
        name for name in aliases
        if targets[name] is not None or actuals[name] is not None
    )
    values = _SignalValues(
        tuple(sorted(set(source_event_ids))),
        targets["calories"], actuals["calories"],
        targets["carbohydrate"], actuals["carbohydrate"],
        targets["protein"], actuals["protein"],
        targets["fat"], actuals["fat"], tolerance,
    )
    if not pairs:
        return _signal(values, "missing", "adherence_evidence_required")
    if any(targets[name] is None or actuals[name] is None for name in pairs):
        return _signal(values, "contradictory", "target_actual_pair_required")
    deviations: list[Decimal] = []
    for name in pairs:
        target_value = targets[name]
        actual_value = actuals[name]
        if target_value is None or actual_value is None:
            raise AdherenceSignalError("complete adherence pair narrowed incorrectly")
        if target_value <= 0:
            return _signal(values, "contradictory", "non_positive_target")
        deviations.append(
            abs(Decimal(actual_value) - Decimal(target_value))
            / Decimal(target_value)
            * Decimal("100")
        )
    status: AdherenceStatus = (
        "adequate" if all(value <= tolerance for value in deviations) else "inadequate"
    )
    return _signal(values, status, None)


build_canonical_adherence_signal = derive_canonical_adherence_signal
adherence_signal_from_target_actual = derive_canonical_adherence_signal
derive_adherence_signal = derive_canonical_adherence_signal
