"""Persisted, digest-bound deterministic reconciliation state."""

from __future__ import annotations

from collections.abc import Mapping
from datetime import date
from typing import Literal

from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator

from checkin_cli.nutrition_onboarding_clarification_policy import (
    POLICY_DIGEST,
    POLICY_VERSION,
    ClarificationPolicyResult,
    compile_clarification_policy,
)
from checkin_cli.nutrition_onboarding_contract import QUESTION_FIELDS, canonical_digest


def _require_digest(value: str, *, field: str) -> str:
    if len(value) != 64 or any(character not in "0123456789abcdef" for character in value):
        raise ValueError(f"{field} must be a lowercase SHA-256 digest")
    return value


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

    issue_id: str
    field: str
    reason_code: str
    question_ko: str = Field(min_length=1, max_length=180)

    @field_validator("issue_id")
    @classmethod
    def require_issue_digest(cls, value: str) -> str:
        return _require_digest(value, field="issue_id")

    @field_validator("field")
    @classmethod
    def require_field(cls, value: str) -> str:
        if value not in QUESTION_FIELDS:
            raise ValueError("issue field is not canonical")
        return value


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

    field: str
    reason_code: Literal["sensitive_input_requires_clinical_review"]


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

    schema_version: Literal["nutrition_input_reconciliation_v2"]
    policy_version: Literal["nutrition_onboarding_clarification_v1"]
    policy_digest: str
    reference_date: date
    compiler_digest: str
    state: Literal["clarifying", "resolved"]
    answers_digest: str
    advisory: dict[str, object]
    issues: tuple[PersistedIssue, ...]
    visible_issue_ids: tuple[str, ...]
    holds: tuple[PersistedHold, ...]
    clarifications: tuple[()] = ()
    current_index: Literal[0] = 0
    digest: str

    @field_validator("policy_digest", "compiler_digest", "answers_digest", "digest")
    @classmethod
    def require_digest(cls, value: str, info: object) -> str:
        return _require_digest(value, field=str(getattr(info, "field_name", "digest")))

    @model_validator(mode="after")
    def validate_record(self) -> "PersistedReconciliation":
        if self.policy_digest != POLICY_DIGEST:
            raise ValueError("reconciliation policy digest is stale")
        expected_state = "clarifying" if self.issues else "resolved"
        if self.state != expected_state:
            raise ValueError("reconciliation state does not match issues")
        issue_ids = tuple(issue.issue_id for issue in self.issues)
        if len(issue_ids) != len(set(issue_ids)):
            raise ValueError("reconciliation issues must be unique")
        if self.visible_issue_ids != issue_ids[:3]:
            raise ValueError("visible issue slice is stale")
        payload = self.model_dump(mode="json", exclude={"digest"})
        if canonical_digest(payload) != self.digest:
            raise ValueError("reconciliation digest is stale")
        return self


def _record_payload(
    result: ClarificationPolicyResult,
    advisory: Mapping[str, object],
) -> dict[str, object]:
    advisory_digest = advisory.get("digest")
    if (
        set(advisory) == {"status", "digest"}
        and advisory.get("status") in {"available", "unavailable"}
        and isinstance(advisory_digest, str)
        and len(advisory_digest) == 64
        and not set(advisory_digest).difference("0123456789abcdef")
    ):
        advisory_payload = dict(advisory)
    else:
        advisory_payload: dict[str, object] = {
            "status": (
                "unavailable"
                if advisory.get("status") == "unavailable"
                else "available"
            ),
            "digest": canonical_digest(dict(advisory)),
        }
    payload: dict[str, object] = {
        "schema_version": "nutrition_input_reconciliation_v2",
        "policy_version": POLICY_VERSION,
        "policy_digest": POLICY_DIGEST,
        "reference_date": result.reference_date.isoformat(),
        "compiler_digest": result.result_digest,
        "state": "clarifying" if result.issues else "resolved",
        "answers_digest": result.answers_digest,
        "advisory": advisory_payload,
        "issues": [issue.model_dump(mode="json") for issue in result.issues],
        "visible_issue_ids": list(result.visible_issue_ids),
        "holds": [hold.model_dump(mode="json") for hold in result.holds],
        "clarifications": [],
        "current_index": 0,
    }
    payload["digest"] = canonical_digest(payload)
    return payload


def build_reconciliation(
    *,
    answers: Mapping[str, object],
    answers_digest: str,
    advisory: Mapping[str, object],
    reference_date: date,
) -> dict[str, object]:
    _require_digest(answers_digest, field="answers_digest")
    result = compile_clarification_policy(
        answers,
        reference_date=reference_date,
    )
    if result.answers_digest != answers_digest:
        raise ValueError("reconciliation answers digest is stale")
    payload = _record_payload(result, advisory)
    return PersistedReconciliation.model_validate(payload).model_dump(mode="json")


def load_reconciliation(value: object) -> dict[str, object]:
    return PersistedReconciliation.model_validate(value).model_dump(mode="json")


def recompute_reconciliation(
    value: object,
    *,
    answers: Mapping[str, object],
    answers_digest: str,
) -> dict[str, object]:
    record = PersistedReconciliation.model_validate(value)
    result = compile_clarification_policy(
        answers,
        reference_date=record.reference_date,
    )
    if result.answers_digest != answers_digest:
        raise ValueError("reconciliation answers digest is stale")
    payload = _record_payload(result, record.advisory)
    return PersistedReconciliation.model_validate(payload).model_dump(mode="json")


def require_resolved_reconciliation(
    value: object,
    *,
    answers_digest: str,
) -> dict[str, object]:
    record = PersistedReconciliation.model_validate(value)
    if record.state != "resolved":
        raise ValueError("input reconciliation is not resolved")
    if record.answers_digest != answers_digest:
        raise ValueError("input reconciliation answers digest is stale")
    return record.model_dump(mode="json")
