"""Experimental immutable cache and compact serialization for terrain profiles."""

from __future__ import annotations

import hashlib
import json
import math
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any

from mountain_twin.terrain.profile import HorizonProfile, HorizonSample, query_profile

CACHE_FORMAT = "terrain_horizon_profile_cache_v0_2"


@dataclass(frozen=True)
class CacheIdentity:
    """Semantic identity for one reusable profile record.

    Coordinates are exact identity fields. Nearby points are never silently
    substituted for one another.
    """

    provider: str
    product: str | None
    source_id: str
    source_sha256: str
    latitude: float
    longitude: float
    surface_semantics: str
    observer_model: str
    observer_height_m: float
    angular_resolution_deg: float
    interpolation_method: str
    horizon_range_m: float
    range_policy: str
    method_version: str

    def __post_init__(self) -> None:
        finite = (
            self.latitude,
            self.longitude,
            self.observer_height_m,
            self.angular_resolution_deg,
            self.horizon_range_m,
        )
        if not all(math.isfinite(value) for value in finite):
            raise ValueError("cache identity numeric fields must be finite")
        if not -90 <= self.latitude <= 90 or not -180 <= self.longitude <= 180:
            raise ValueError("cache identity coordinates are invalid")
        if self.observer_height_m < 0 or self.angular_resolution_deg <= 0:
            raise ValueError("cache identity height/resolution are invalid")
        if self.horizon_range_m <= 0:
            raise ValueError("cache identity horizon range must be positive")
        if self.interpolation_method not in {"circular_linear", "nearest"}:
            raise ValueError("unsupported cache interpolation method")

    def manifest(self) -> dict[str, Any]:
        return asdict(self)

    def key(self) -> str:
        encoded = json.dumps(
            self.manifest(), sort_keys=True, separators=(",", ":"), allow_nan=False
        ).encode("utf-8")
        return hashlib.sha256(encoded).hexdigest()


class CacheSemanticMismatch(ValueError):
    """Raised when a stored profile does not match the requested identity."""


def _profile_manifest(profile: HorizonProfile) -> dict[str, Any]:
    value = profile.to_dict()
    value.pop("samples")
    return value


def serialize_profile(profile: HorizonProfile, identity: CacheIdentity) -> dict[str, Any]:
    """Return a deterministic compact JSON-compatible cache envelope."""
    _validate_profile_identity(profile, identity)
    states = sorted({sample.convergence_state for sample in profile.samples})
    reasons = sorted({code for sample in profile.samples for code in sample.reason_codes})
    state_codes = {state: index for index, state in enumerate(states)}
    reason_codes = {code: index for index, code in enumerate(reasons)}
    return {
        "format": CACHE_FORMAT,
        "identity": identity.manifest(),
        "profile": _profile_manifest(profile),
        "samples": {
            "azimuth_deg": [sample.azimuth_deg for sample in profile.samples],
            "horizon_angle_deg": [sample.horizon_angle_deg for sample in profile.samples],
            "controlling_distance_m": [sample.controlling_distance_m for sample in profile.samples],
            "final_range_m": [sample.final_range_m for sample in profile.samples],
            "convergence_state_table": states,
            "convergence_state": [
                state_codes[sample.convergence_state] for sample in profile.samples
            ],
            "reason_code_table": reasons,
            "reason_codes": [
                [reason_codes[code] for code in sample.reason_codes] for sample in profile.samples
            ],
        },
    }


def profile_from_dict(value: dict[str, Any]) -> HorizonProfile:
    """Restore the verbose v0.1 profile representation for cache migration."""
    data = dict(value)
    samples = tuple(
        HorizonSample(
            **{
                **sample,
                "reason_codes": tuple(sample.get("reason_codes", ())),
            }
        )
        for sample in data.pop("samples")
    )
    data["limitations"] = tuple(data.get("limitations", ()))
    return HorizonProfile(samples=samples, **data)


def deserialize_profile(payload: dict[str, Any]) -> tuple[CacheIdentity, HorizonProfile]:
    """Restore a cache envelope and preserve all sample semantic states."""
    if payload.get("format") != CACHE_FORMAT:
        raise ValueError("unsupported horizon profile cache format")
    identity = CacheIdentity(**payload["identity"])
    profile_data = dict(payload["profile"])
    profile_data["limitations"] = tuple(profile_data.get("limitations", ()))
    samples_data = payload["samples"]
    lengths = {
        len(samples_data[name])
        for name in (
            "azimuth_deg",
            "horizon_angle_deg",
            "controlling_distance_m",
            "final_range_m",
            "convergence_state",
            "reason_codes",
        )
    }
    if len(lengths) != 1:
        raise ValueError("profile sample arrays have inconsistent lengths")
    states = samples_data["convergence_state_table"]
    reasons = samples_data["reason_code_table"]
    samples = tuple(
        HorizonSample(
            azimuth_deg=azimuth,
            horizon_angle_deg=horizon,
            controlling_distance_m=distance,
            final_range_m=final_range,
            convergence_state=states[state_code],
            reason_codes=tuple(reasons[code] for code in reason_code_list),
        )
        for azimuth, horizon, distance, final_range, state_code, reason_code_list in zip(
            samples_data["azimuth_deg"],
            samples_data["horizon_angle_deg"],
            samples_data["controlling_distance_m"],
            samples_data["final_range_m"],
            samples_data["convergence_state"],
            samples_data["reason_codes"],
        )
    )
    profile = HorizonProfile(samples=samples, **profile_data)
    _validate_profile_identity(profile, identity)
    return identity, profile


def _validate_profile_identity(profile: HorizonProfile, identity: CacheIdentity) -> None:
    pairs = (
        (profile.provider, identity.provider, "provider"),
        (profile.product, identity.product, "product"),
        (profile.latitude, identity.latitude, "latitude"),
        (profile.longitude, identity.longitude, "longitude"),
        (profile.surface_semantics, identity.surface_semantics, "surface_semantics"),
        (profile.observer_model, identity.observer_model, "observer_model"),
        (profile.observer_height_m, identity.observer_height_m, "observer_height_m"),
        (profile.azimuth_resolution_deg, identity.angular_resolution_deg, "resolution"),
        (profile.interpolation_policy, identity.interpolation_method, "interpolation"),
    )
    for actual, expected, field in pairs:
        if actual != expected:
            raise CacheSemanticMismatch(f"profile identity mismatch: {field}")


class HorizonProfileCache:
    """Filesystem-backed immutable cache with semantic lookup verification."""

    def __init__(self, root: Path):
        self.root = Path(root)

    def _path(self, identity: CacheIdentity) -> Path:
        return self.root / f"{identity.key()}.json"

    def put_profile(self, identity: CacheIdentity, profile: HorizonProfile) -> str:
        payload = serialize_profile(profile, identity)
        text = json.dumps(payload, sort_keys=True, separators=(",", ":"), allow_nan=False) + "\n"
        path = self._path(identity)
        self.root.mkdir(parents=True, exist_ok=True)
        if path.exists():
            if path.read_text(encoding="utf-8") != text:
                raise CacheSemanticMismatch("immutable cache key already has different content")
        else:
            path.write_text(text, encoding="utf-8")
        return identity.key()

    def get_profile(self, identity: CacheIdentity) -> HorizonProfile:
        path = self._path(identity)
        payload = json.loads(path.read_text(encoding="utf-8"))
        stored_identity, profile = deserialize_profile(payload)
        if stored_identity != identity:
            raise CacheSemanticMismatch("stored cache identity differs from requested identity")
        return profile

    def has_profile(self, identity: CacheIdentity) -> bool:
        try:
            self.get_profile(identity)
        except (FileNotFoundError, json.JSONDecodeError, KeyError, TypeError, ValueError):
            return False
        return True

    def query_horizon(self, identity: CacheIdentity, azimuth_deg: float):
        return query_profile(
            self.get_profile(identity), azimuth_deg, method=identity.interpolation_method
        )
