"""Deterministic clarification and safety policy for nutrition onboarding."""

from __future__ import annotations

from collections.abc import Mapping
from copy import deepcopy
from datetime import date
from decimal import Decimal, InvalidOperation
from typing import Literal, cast

from pydantic import BaseModel, ConfigDict

from checkin_cli.nutrition_onboarding_clarification_rules import (
    POLICY_DIGEST,
    POLICY_VERSION,
    QUESTION_COPY as _QUESTION_COPY,
    SENSITIVE_FIELDS as _SENSITIVE_FIELDS,
    VISIBLE_ISSUE_LIMIT as _VISIBLE_ISSUE_LIMIT,
)
from checkin_cli.nutrition_onboarding_contract import QUESTION_FIELDS, canonical_digest
from checkin_cli.nutrition_onboarding_models import (
    build_onboarding_baseline,
    canonical_equation_sex_basis,
    derive_adult_age,
)

class _PolicyModel(BaseModel):
    model_config = ConfigDict(extra="forbid", frozen=True)


class ClarificationIssue(_PolicyModel):
    issue_id: str
    field: str
    reason_code: str
    question_ko: str


class SafetyHold(_PolicyModel):
    field: str
    reason_code: Literal["sensitive_input_requires_clinical_review"]


class ClarificationPolicyResult(_PolicyModel):
    policy_version: str
    policy_digest: str
    reference_date: date
    answers: dict[str, object]
    answers_digest: str
    issues: tuple[ClarificationIssue, ...]
    visible_issue_ids: tuple[str, ...]
    holds: tuple[SafetyHold, ...]
    result_digest: str


def _canonical_decimal(value: object) -> str | None:
    if isinstance(value, bool) or not isinstance(value, (str, int, float, Decimal)):
        return None
    parsed = _decimal(value)
    if parsed is None or not parsed.is_finite():
        return None
    rendered = format(parsed.normalize(), "f")
    return rendered.rstrip("0").rstrip(".") if "." in rendered else rendered


def canonicalize_answer(field: str, value: object) -> object:
    """Canonicalize one boundary value without accepting an invalid type."""
    if field == "equation_sex_basis":
        return canonical_equation_sex_basis(value) or value
    if field in {"height_cm", "weight_kg", "target_weight_kg"} and value is not None:
        return _canonical_decimal(value) or value
    return deepcopy(value)


def _canonical_answers(answers: Mapping[str, object]) -> dict[str, object]:
    return {
        field: canonicalize_answer(field, answers[field])
        for field in QUESTION_FIELDS
        if field in answers
    }


def _is_ambiguous(value: object) -> bool:
    if isinstance(value, str):
        return any(marker in value for marker in (" 또는 ", " 혹은 "))
    if isinstance(value, Mapping):
        return any(_is_ambiguous(item) for item in value.values())
    if isinstance(value, (list, tuple)):
        return any(_is_ambiguous(item) for item in value)
    return False


def _valid_structured(value: object) -> bool:
    if not isinstance(value, dict) or set(value) != {"status", "items"}:
        return False
    structured = cast(dict[str, object], value)
    status = structured.get("status")
    items = structured.get("items")
    if not isinstance(items, list) or any(not isinstance(item, str) for item in items):
        return False
    typed_items = cast(list[str], items)
    return (status == "none" and not typed_items) or (
        status == "provided"
        and bool(typed_items)
        and all(item.strip() for item in typed_items)
    )


def _sensitive(field: str, value: object) -> bool:
    if field in {"conditions", "medications"}:
        return _valid_structured(value) and cast(
            dict[str, object], value
        ).get("status") == "provided"
    return (value is True or value is None) and field in {
        "pregnancy_breastfeeding",
        "eating_disorder_risk",
    }


def _safe_adult_dob(reference_date: date) -> str:
    try:
        return reference_date.replace(year=reference_date.year - 30).isoformat()
    except ValueError:
        return reference_date.replace(year=reference_date.year - 30, day=28).isoformat()


def _decimal(value: object) -> Decimal | None:
    try:
        return Decimal(str(value))
    except (InvalidOperation, TypeError, ValueError):
        return None


def _type_reasons(answers: Mapping[str, object], reasons: dict[str, str]) -> None:
    structured = {
        "allergies",
        "intolerances",
        "religious_ethical_exclusions",
        "disliked_foods",
        "dietary_preferences",
        "conditions",
        "medications",
    }
    for field in structured:
        if field in answers and not _valid_structured(answers[field]):
            reasons.setdefault(field, "invalid_value")
    for field in ("pregnancy_breastfeeding", "eating_disorder_risk"):
        if field in answers and answers[field] is not None and type(answers[field]) is not bool:
            reasons.setdefault(field, "invalid_value")
    if "meal_count" in answers and type(answers["meal_count"]) is not int:
        reasons.setdefault("meal_count", "invalid_value")


def _cross_field_reasons(answers: Mapping[str, object], reasons: dict[str, str]) -> None:
    goal = answers.get("goal_type")
    target_weight = answers.get("target_weight_kg")
    target_date = answers.get("target_date")
    if goal == "maintain":
        if target_weight is not None:
            reasons.setdefault("target_weight_kg", "target_not_allowed")
        if target_date is not None:
            reasons.setdefault("target_date", "target_not_allowed")
        return
    if goal not in {"loss", "gain"}:
        return
    if target_weight is None:
        reasons.setdefault("target_weight_kg", "target_required")
    if target_date is None:
        reasons.setdefault("target_date", "target_required")
    current = _decimal(answers.get("weight_kg"))
    target = _decimal(target_weight)
    if current is None or target is None:
        return
    if goal == "loss" and target >= current:
        reasons.setdefault("target_weight_kg", "target_direction_invalid")
    if goal == "gain" and target <= current:
        reasons.setdefault("target_weight_kg", "target_direction_invalid")


def _baseline_reasons(
    answers: Mapping[str, object],
    reasons: dict[str, str],
    reference_date: date,
) -> None:
    payload = dict(answers)
    raw_dob = payload.get("date_of_birth")
    try:
        if not isinstance(raw_dob, str):
            raise ValueError("date_of_birth is required")
        parsed_dob = date.fromisoformat(raw_dob)
        derive_adult_age(parsed_dob, as_of=reference_date)
    except (TypeError, ValueError):
        reasons.setdefault("date_of_birth", "invalid_date_of_birth")
        payload["date_of_birth"] = _safe_adult_dob(reference_date)
    payload.update(
        schema_version="1.0",
        customer_key="policy_validation",
        session_status="completed",
        review_decision="pending",
        publication_status="unpublished",
    )
    try:
        build_onboarding_baseline(payload, as_of=reference_date)
    except ValueError as exc:
        errors = getattr(exc, "errors", lambda: ())()
        for error in errors:
            location = error.get("loc", ())
            if not location:
                continue  # Explicit cross-field rules above own root errors.
            field = str(location[0])
            if field == "adult_age":
                field = "date_of_birth"
            if field in QUESTION_FIELDS:
                reasons.setdefault(field, "invalid_value")


def _issue(field: str, reason_code: str, value: object) -> ClarificationIssue:
    return ClarificationIssue(
        issue_id=canonical_digest(
            {
                "policy_version": POLICY_VERSION,
                "policy_digest": POLICY_DIGEST,
                "field": field,
                "reason_code": reason_code,
                "value_digest": canonical_digest(value),
            }
        ),
        field=field,
        reason_code=reason_code,
        question_ko=_QUESTION_COPY[field],
    )


def compile_clarification_policy(
    answers: Mapping[str, object],
    *,
    reference_date: date,
) -> ClarificationPolicyResult:
    """Compile every ordered issue and hold against an explicit reference date."""
    canonical = _canonical_answers(answers)
    answers_digest = canonical_digest(canonical)
    reasons: dict[str, str] = {}
    for field in QUESTION_FIELDS:
        if field not in canonical:
            reasons[field] = "missing_value"
        elif _is_ambiguous(canonical[field]):
            reasons[field] = (
                "multiple_measurements"
                if field in {"height_cm", "weight_kg", "target_weight_kg"}
                else "multiple_values"
            )
    _type_reasons(canonical, reasons)
    _cross_field_reasons(canonical, reasons)
    _baseline_reasons(canonical, reasons, reference_date)

    holds = tuple(
        SafetyHold(field=field, reason_code="sensitive_input_requires_clinical_review")
        for field in _SENSITIVE_FIELDS
        if field in canonical and _sensitive(field, canonical[field])
    )
    for hold in holds:
        reasons.pop(hold.field, None)
    issues = tuple(
        _issue(field, reasons[field], canonical.get(field))
        for field in QUESTION_FIELDS
        if field in reasons
    )
    visible_issue_ids = tuple(issue.issue_id for issue in issues[:_VISIBLE_ISSUE_LIMIT])
    result_payload = {
        "policy_version": POLICY_VERSION,
        "policy_digest": POLICY_DIGEST,
        "reference_date": reference_date.isoformat(),
        "answers_digest": answers_digest,
        "issues": [issue.model_dump(mode="json") for issue in issues],
        "visible_issue_ids": visible_issue_ids,
        "holds": [hold.model_dump(mode="json") for hold in holds],
    }
    return ClarificationPolicyResult(
        policy_version=POLICY_VERSION,
        policy_digest=POLICY_DIGEST,
        reference_date=reference_date,
        answers=canonical,
        answers_digest=answers_digest,
        issues=issues,
        visible_issue_ids=visible_issue_ids,
        holds=holds,
        result_digest=canonical_digest(result_payload),
    )
