"""RED contracts for deterministic initial nutrition calculations."""

from __future__ import annotations

import importlib.util
from datetime import date
from decimal import Decimal

import pytest
from pydantic import ValidationError


MODULE = "checkin_cli.nutrition_onboarding_calculations"
assert importlib.util.find_spec(MODULE) is not None, (
    "nutrition onboarding contract missing: "
    "checkin_cli.nutrition_onboarding_calculations"
)

from checkin_cli.nutrition_onboarding_calculations import (  # noqa: E402
    CalculationDeclinedError,
    InitialNutritionPlan,
    UnsafeGoalTrajectoryError,
    WeeklyNutritionTarget,
    activity_factor,
    generate_initial_plan,
    mifflin_st_jeor_bmr,
)
from checkin_cli.nutrition_onboarding_models import (  # noqa: E402
    ActivityCategory,
    EquationSexBasis,
    GoalType,
    NutritionOnboardingBaseline,
    OnboardingSessionStatus,
    PublicationStatus,
    ReviewDecision,
    StructuredItemStatus,
    StructuredItems,
)


def _none() -> StructuredItems:
    return StructuredItems(status=StructuredItemStatus.NONE, items=())


def _baseline(**overrides: object) -> NutritionOnboardingBaseline:
    values: dict[str, object] = {
        "schema_version": "1.0",
        "customer_key": "client_001",
        "adult_age": 30,
        "equation_sex_basis": EquationSexBasis.MALE,
        "height_cm": Decimal("180"),
        "weight_kg": Decimal("80"),
        "activity_category": ActivityCategory.MODERATE,
        "goal_type": GoalType.MAINTAIN,
        "target_weight_kg": None,
        "target_date": None,
        "dietary_preferences": _none(),
        "disliked_foods": _none(),
        "allergies": _none(),
        "intolerances": _none(),
        "religious_ethical_exclusions": _none(),
        "conditions": _none(),
        "medications": _none(),
        "session_status": OnboardingSessionStatus.COMPLETED,
        "review_decision": ReviewDecision.APPROVED,
        "publication_status": PublicationStatus.UNPUBLISHED,
    }
    values.update(overrides)
    return NutritionOnboardingBaseline(**values)


def test_mifflin_st_jeor_golden_vectors_and_decline_path() -> None:
    assert mifflin_st_jeor_bmr(
        weight_kg=Decimal("80"),
        height_cm=Decimal("180"),
        adult_age=30,
        equation_sex_basis=EquationSexBasis.MALE,
    ) == Decimal("1780")
    assert mifflin_st_jeor_bmr(
        weight_kg=Decimal("80"),
        height_cm=Decimal("180"),
        adult_age=30,
        equation_sex_basis=EquationSexBasis.FEMALE,
    ) == Decimal("1614")
    assert mifflin_st_jeor_bmr(
        weight_kg=Decimal("80"),
        height_cm=Decimal("180"),
        adult_age=30,
        equation_sex_basis=EquationSexBasis.DECLINE,
    ) is None


def test_activity_factors_are_pinned_decimal_constants() -> None:
    assert {
        category: activity_factor(category)
        for category in ActivityCategory
    } == {
        ActivityCategory.SEDENTARY: Decimal("1.2"),
        ActivityCategory.LIGHT: Decimal("1.375"),
        ActivityCategory.MODERATE: Decimal("1.55"),
        ActivityCategory.VERY_ACTIVE: Decimal("1.725"),
        ActivityCategory.EXTRA_ACTIVE: Decimal("1.9"),
    }


def test_maintain_plan_matches_the_independent_golden_vector() -> None:
    plan = generate_initial_plan(
        _baseline(),
        starts_on=date(2026, 8, 3),
    )

    assert plan.method_id == "mifflin_st_jeor_1990"
    assert plan.method_version == "1.0"
    assert plan.bmr_kcal == Decimal("1780")
    assert plan.tdee_kcal == Decimal("2759")
    assert plan.starts_on == date(2026, 8, 3)
    assert len(plan.weeks) == 12
    assert tuple(row.week for row in plan.weeks) == tuple(range(1, 13))
    assert [
        (
            row.week,
            row.calories_kcal,
            row.protein_g,
            row.carbohydrate_g,
            row.fat_g,
        )
        for row in plan.weeks
    ] == [
        (week, 2760, 160, 350, 80)
        for week in range(1, 13)
    ]


def test_every_week_reconciles_calories_exactly_with_existing_solver_semantics() -> None:
    plan = generate_initial_plan(
        _baseline(),
        starts_on=date(2026, 8, 3),
    )

    for row in plan.weeks:
        assert row.calories_kcal == (
            4 * row.protein_g
            + 4 * row.carbohydrate_g
            + 9 * row.fat_g
        )


def test_plan_and_week_schemas_are_frozen_and_forbid_extra_fields() -> None:
    plan = generate_initial_plan(
        _baseline(),
        starts_on=date(2026, 8, 3),
    )
    row = plan.weeks[0]

    with pytest.raises(ValidationError, match="Extra inputs are not permitted"):
        WeeklyNutritionTarget(**row.model_dump(), food="chicken")
    with pytest.raises(ValidationError, match="Extra inputs are not permitted"):
        InitialNutritionPlan(
            **plan.model_dump(),
            recommendation="eat oats",
        )
    with pytest.raises(ValidationError, match="frozen"):
        row.calories_kcal = 2000
    with pytest.raises(ValidationError, match="frozen"):
        plan.weeks = ()


def test_plan_requires_exactly_twelve_sequential_rows() -> None:
    plan = generate_initial_plan(
        _baseline(),
        starts_on=date(2026, 8, 3),
    )

    with pytest.raises(ValidationError, match="12"):
        InitialNutritionPlan(
            **{
                **plan.model_dump(),
                "weeks": plan.weeks[:-1],
            }
        )
    with pytest.raises(ValidationError, match="1 through 12"):
        InitialNutritionPlan(
            **{
                **plan.model_dump(),
                "weeks": (
                    plan.weeks[0].model_copy(update={"week": 2}),
                    *plan.weeks[1:],
                ),
            }
        )


@pytest.mark.parametrize(
    ("field", "value"),
    (
        ("calories_kcal", 1499),
        ("calories_kcal", 4501),
        ("protein_g", 119),
        ("protein_g", 251),
        ("fat_g", 39),
        ("fat_g", 151),
        ("carbohydrate_g", -1),
    ),
)
def test_week_targets_enforce_global_calorie_and_macro_bounds(
    field: str,
    value: int,
) -> None:
    valid: dict[str, object] = {
        "week": 1,
        "calories_kcal": 2300,
        "protein_g": 150,
        "carbohydrate_g": 290,
        "fat_g": 60,
    }
    valid[field] = value

    with pytest.raises(ValidationError):
        WeeklyNutritionTarget(**valid)


def test_week_target_accepts_closed_global_boundaries_when_reconciled() -> None:
    low = WeeklyNutritionTarget(
        week=1,
        calories_kcal=1500,
        protein_g=120,
        carbohydrate_g=165,
        fat_g=40,
    )
    high = WeeklyNutritionTarget(
        week=12,
        calories_kcal=4500,
        protein_g=250,
        carbohydrate_g=542,
        fat_g=148,
    )
    max_fat = WeeklyNutritionTarget(
        week=12,
        calories_kcal=4498,
        protein_g=250,
        carbohydrate_g=537,
        fat_g=150,
    )

    assert low.calories_kcal == 4 * 120 + 4 * 165 + 9 * 40
    assert high.calories_kcal == 4 * 250 + 4 * 542 + 9 * 148
    assert max_fat.calories_kcal == 4 * 250 + 4 * 537 + 9 * 150


def test_week_target_rejects_non_reconciling_macro_rows() -> None:
    with pytest.raises(ValidationError, match="reconcile"):
        WeeklyNutritionTarget(
            week=1,
            calories_kcal=2301,
            protein_g=150,
            carbohydrate_g=290,
            fat_g=60,
        )


@pytest.mark.parametrize(
    ("baseline", "limit"),
    (
        (
            _baseline(
                goal_type=GoalType.LOSS,
                target_weight_kg=Decimal("70"),
                target_date=date(2026, 8, 31),
            ),
            Decimal("0.01"),
        ),
        (
            _baseline(
                goal_type=GoalType.GAIN,
                target_weight_kg=Decimal("90"),
                target_date=date(2026, 8, 31),
            ),
            Decimal("0.005"),
        ),
    ),
    ids=("unsafe-loss", "unsafe-gain"),
)
def test_requested_goal_is_kept_with_safe_date_recommendation(
    baseline: NutritionOnboardingBaseline,
    limit: Decimal,
) -> None:
    plan = generate_initial_plan(
        baseline,
        starts_on=date(2026, 8, 3),
    )

    assert plan.requested_trajectory_within_guardrail is False
    assert plan.recommended_target_date is not None
    assert plan.recommended_target_date > baseline.target_date
    first_week_change = plan.projected_weights_kg[1] - baseline.weight_kg
    assert abs(first_week_change / baseline.weight_kg) <= limit


def test_near_boundary_loss_is_capped_to_exact_safe_rate() -> None:
    baseline = _baseline(
        weight_kg=Decimal("82.3"),
        goal_type=GoalType.LOSS,
        target_weight_kg=Decimal("75"),
        target_date=date(2026, 10, 22),
    )

    plan = generate_initial_plan(
        baseline,
        starts_on=date(2026, 8, 21),
    )

    assert plan.requested_trajectory_within_guardrail is False
    assert plan.recommended_target_date == date(2026, 10, 23)
    first_week_change = plan.projected_weights_kg[1] - baseline.weight_kg
    assert abs(first_week_change / baseline.weight_kg) <= Decimal("0.01")


def test_declined_equation_basis_does_not_guess_a_plan() -> None:
    with pytest.raises(
        CalculationDeclinedError,
        match="equation sex basis",
    ):
        generate_initial_plan(
            _baseline(
                equation_sex_basis=EquationSexBasis.DECLINE,
            ),
            starts_on=date(2026, 8, 3),
        )


def test_calculation_output_contains_targets_not_food_recommendations() -> None:
    plan = generate_initial_plan(
        _baseline(),
        starts_on=date(2026, 8, 3),
    )
    payload = plan.model_dump(mode="json")
    forbidden_fragments = ("food", "meal", "recipe", "supplement")

    def keys(value: object) -> list[str]:
        if isinstance(value, dict):
            return [str(key) for key in value] + [
                nested
                for child in value.values()
                for nested in keys(child)
            ]
        if isinstance(value, list):
            return [
                nested
                for child in value
                for nested in keys(child)
            ]
        return []

    assert set(payload["weeks"][0]) == {
        "week",
        "calories_kcal",
        "protein_g",
        "carbohydrate_g",
        "fat_g",
    }
    assert not any(
        fragment in key.lower()
        for key in keys(payload)
        for fragment in forbidden_fragments
    )
