"""Transparent route interpolation and elevation-context diagnostics."""

from __future__ import annotations

import json
import math
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import Any, Iterable, Sequence

from mountain_twin.weather.provider import HOURLY_VARIABLES, weather_variable_semantics

SCALAR_VARIABLES = frozenset(
    {
        "temperature_2m",
        "relative_humidity_2m",
        "apparent_temperature",
        "cloud_cover",
        "cloud_cover_low",
        "cloud_cover_mid",
        "cloud_cover_high",
        "wind_speed_10m",
        "wind_gusts_10m",
        "visibility",
        "freezing_level_height",
    }
)
NEAREST_CONTEXT_VARIABLES = frozenset(
    {
        "precipitation",
        "rain",
        "snowfall",
        "snow_depth",
        "precipitation_probability",
        "weather_code",
    }
)
CATEGORICAL_VARIABLES = frozenset({"weather_code"})


@dataclass(frozen=True)
class RouteWeatherSample:
    sample_id: str
    point_index: int
    route_distance_m: float
    latitude: float
    longitude: float
    route_elevation_m: float | None
    provider_elevation_m: float | None
    elevation_difference_m: float | None
    selection_reasons: tuple[str, ...]
    variables: dict[str, float | int | None]
    units: dict[str, str | None]
    missing_variable_reasons: dict[str, str] = field(default_factory=dict)
    variable_semantics: dict[str, dict[str, Any]] = field(default_factory=dict)


@dataclass(frozen=True)
class WeatherField:
    source_type: str
    scenario_datetime: str
    timezone: str
    provider: str
    model_selection: str
    samples: tuple[RouteWeatherSample, ...]


@dataclass(frozen=True)
class RouteWeatherValue:
    value: float | int | None
    representation_state: str
    method: str
    contributing_sample_ids: tuple[str, ...]
    support_distances_m: tuple[float, ...]
    elevation_adjustment_state: str = "not_applied"
    limitation_codes: tuple[str, ...] = ()


def load_weather_field(path: Path) -> WeatherField:
    """Load one committed v0.1 compact artifact without contacting a provider."""
    document = json.loads(path.read_text(encoding="utf-8"))
    samples = []
    for item in document["samples"]:
        selection = item["selection"]
        contract = item["contract"]
        payload = contract["payload"]
        samples.append(
            RouteWeatherSample(
                sample_id=selection["sample_id"],
                point_index=selection["point_index"],
                route_distance_m=selection["route_distance_m"],
                latitude=selection["latitude"],
                longitude=selection["longitude"],
                route_elevation_m=selection["route_elevation_m"],
                provider_elevation_m=payload["provider_grid_elevation_m"],
                elevation_difference_m=payload["route_minus_provider_elevation_m"],
                selection_reasons=tuple(selection["selection_reasons"]),
                variables=dict(payload["variables"]),
                units=dict(payload["units"]),
                missing_variable_reasons=_missing_variable_reasons(payload),
                variable_semantics=_variable_semantics(payload),
            )
        )
    ordered = tuple(sorted(samples, key=lambda item: item.route_distance_m))
    if [item.sample_id for item in samples] != [item.sample_id for item in ordered]:
        raise ValueError("weather samples must be ordered by route progression")
    return WeatherField(
        source_type=document["source_type"],
        scenario_datetime=document["scenario_datetime"],
        timezone=document["timezone"],
        provider=document["provider"],
        model_selection=document["model_selection"],
        samples=ordered,
    )


def interpolate_value(
    samples: Sequence[RouteWeatherSample],
    route_distance_m: float,
    variable: str,
    *,
    method: str = "policy",
) -> RouteWeatherValue:
    """Interpolate one variable while retaining its spatial support state."""
    if variable not in HOURLY_VARIABLES:
        raise ValueError(f"unsupported weather variable: {variable}")
    if not samples:
        return _unknown("WEATHER_NO_SAMPLES")
    exact = _exact_sample(samples, route_distance_m)
    if exact is not None:
        return _sample_value(exact, variable)
    if (
        route_distance_m < samples[0].route_distance_m
        or route_distance_m > samples[-1].route_distance_m
    ):
        return _unknown("WEATHER_ROUTE_EXTRAPOLATION_UNSUPPORTED")
    left, right = _bracket(samples, route_distance_m)
    if method == "nearest" or (method == "policy" and variable in NEAREST_CONTEXT_VARIABLES):
        selected = min(
            (left, right), key=lambda item: abs(item.route_distance_m - route_distance_m)
        )
        if selected.variables.get(variable) is None:
            return RouteWeatherValue(
                None,
                "UNKNOWN",
                "nearest_route_sample",
                (selected.sample_id,),
                (abs(selected.route_distance_m - route_distance_m),),
                limitation_codes=(
                    selected.missing_variable_reasons.get(
                        variable, "WEATHER_NEAREST_SAMPLE_MISSING"
                    ),
                ),
            )
        limitation = (
            ("WEATHER_NEAREST_CONTEXT_NOT_SPATIAL_RECONSTRUCTION",) if method == "policy" else ()
        )
        return RouteWeatherValue(
            selected.variables.get(variable),
            "NEAREST_CONTEXT",
            "nearest_route_sample",
            (selected.sample_id,),
            (abs(selected.route_distance_m - route_distance_m),),
            limitation_codes=limitation,
        )
    if method not in {"linear", "policy"}:
        raise ValueError(f"unknown weather interpolation method: {method}")
    first, second = left.variables.get(variable), right.variables.get(variable)
    if first is None or second is None:
        return _unknown("WEATHER_INTERPOLATION_INPUT_MISSING")
    fraction = (route_distance_m - left.route_distance_m) / (
        right.route_distance_m - left.route_distance_m
    )
    if variable == "wind_direction_10m":
        value = _circular_interpolate(float(first), float(second), fraction)
        interpolation = "circular_vector_linear_route_distance"
    elif variable in CATEGORICAL_VARIABLES:
        return _unknown("WEATHER_CATEGORICAL_INTERPOLATION_UNDEFINED")
    elif variable in NEAREST_CONTEXT_VARIABLES:
        value = float(first) + fraction * (float(second) - float(first))
        interpolation = "linear_route_distance_context_only"
    elif variable in SCALAR_VARIABLES:
        value = float(first) + fraction * (float(second) - float(first))
        interpolation = "linear_route_distance"
    else:
        return _unknown("WEATHER_VARIABLE_INTERPOLATION_UNDEFINED")
    limitations = ["WEATHER_ROUTE_DISTANCE_INTERPOLATION_IS_NOT_GRID_RECONSTRUCTION"]
    if variable in NEAREST_CONTEXT_VARIABLES:
        limitations.append("WEATHER_VARIABLE_INTERPOLATION_DOES_NOT_RECONSTRUCT_SPATIAL_CELLS")
    return RouteWeatherValue(
        value,
        "INTERPOLATED",
        interpolation,
        (left.sample_id, right.sample_id),
        (route_distance_m - left.route_distance_m, right.route_distance_m - route_distance_m),
        limitation_codes=tuple(limitations),
    )


def route_representation(
    field: WeatherField,
    route_points: Sequence[Any],
    *,
    method: str = "policy",
) -> list[dict[str, Any]]:
    """Build an experimental full-route representation; never replaces v0.1 data."""
    rows = []
    for point in route_points:
        route_elevation = getattr(point, "elevation_m", getattr(point, "gpx_elevation_m", None))
        route_distance = getattr(
            point, "route_distance_m", getattr(point, "cumulative_distance_m", None)
        )
        row = {
            "route_id": point.route_id,
            "point_index": point.point_index,
            "latitude": point.latitude,
            "longitude": point.longitude,
            "route_distance_m": route_distance,
            "route_elevation_m": route_elevation,
            "scenario_datetime": field.scenario_datetime,
            "timezone": field.timezone,
            "provider": field.provider,
            "source_type": field.source_type,
            "model_selection": field.model_selection,
            "representation_method": method,
            "values": {
                variable: asdict(
                    interpolate_value(field.samples, route_distance, variable, method=method)
                )
                for variable in HOURLY_VARIABLES
            },
        }
        rows.append(row)
    return rows


def elevation_sensitivity(
    field: WeatherField,
    lapse_rates_c_per_km: Sequence[float] = (6.5, 4.0, 9.8),
) -> dict[str, Any]:
    """Compare illustrative lapse-rate scenarios; no correction is selected."""
    records = []
    for sample in field.samples:
        delta = sample.elevation_difference_m
        adjustments = {
            str(rate): None if delta is None else -rate * delta / 1000
            for rate in lapse_rates_c_per_km
        }
        records.append(
            {
                "sample_id": sample.sample_id,
                "route_minus_provider_elevation_m": delta,
                "raw_temperature_2m": sample.variables.get("temperature_2m"),
                "adjustments_c": adjustments,
                "adjusted_temperature_2m": {
                    str(rate): None
                    if delta is None or sample.variables.get("temperature_2m") is None
                    else sample.variables["temperature_2m"] + adjustments[str(rate)]
                    for rate in lapse_rates_c_per_km
                },
            }
        )
    summary = {"lapse_rates_c_per_km": list(lapse_rates_c_per_km), "samples": records}
    for rate in lapse_rates_c_per_km:
        values = [
            item["adjustments_c"][str(rate)]
            for item in records
            if item["adjustments_c"][str(rate)] is not None
        ]
        temps = [
            item["adjusted_temperature_2m"][str(rate)]
            for item in records
            if item["adjusted_temperature_2m"][str(rate)] is not None
        ]
        summary[str(rate)] = {
            "adjustment_min_c": min(values) if values else None,
            "adjustment_median_c": _median(values),
            "adjustment_max_c": max(values) if values else None,
            "adjusted_temperature_min_c": min(temps) if temps else None,
            "adjusted_temperature_max_c": max(temps) if temps else None,
        }
    return summary


def leave_one_out(
    field: WeatherField,
    variables: Iterable[str],
) -> dict[str, Any]:
    """Compare neighboring-sample estimates with the withheld provider field."""
    result = {}
    for variable in variables:
        errors_linear, errors_nearest = [], []
        direction_errors = []
        for index, sample in enumerate(field.samples[1:-1], start=1):
            retained = field.samples[:index] + field.samples[index + 1 :]
            linear = interpolate_value(retained, sample.route_distance_m, variable, method="linear")
            nearest = interpolate_value(
                retained, sample.route_distance_m, variable, method="nearest"
            )
            actual = sample.variables.get(variable)
            if actual is None or linear.value is None or nearest.value is None:
                continue
            if variable == "wind_direction_10m":
                errors_linear.append(_angular_error(float(linear.value), float(actual)))
                errors_nearest.append(_angular_error(float(nearest.value), float(actual)))
                direction_errors.append(errors_linear[-1])
            else:
                errors_linear.append(abs(float(linear.value) - float(actual)))
                errors_nearest.append(abs(float(nearest.value) - float(actual)))
        result[variable] = {
            "sample_count": len(errors_linear),
            "linear": _error_summary(errors_linear),
            "nearest": _error_summary(errors_nearest),
            "error_semantics": "agreement_with_sampled_provider_field_not_atmospheric_observation",
        }
    return result


def compare_fields(first: WeatherField, second: WeatherField) -> dict[str, Any]:
    result = {}
    for variable in ("temperature_2m", "cloud_cover", "wind_speed_10m", "wind_gusts_10m"):
        first_values = [item.variables.get(variable) for item in first.samples]
        second_values = [item.variables.get(variable) for item in second.samples]
        result[variable] = {
            "forecast_range": [min(first_values), max(first_values)],
            "historical_forecast_range": [min(second_values), max(second_values)],
            "forecast_adjacent_abs_change_mean": _adjacent_mean(first_values),
            "historical_forecast_adjacent_abs_change_mean": _adjacent_mean(second_values),
        }
    return result


def _exact_sample(
    samples: Sequence[RouteWeatherSample], distance: float
) -> RouteWeatherSample | None:
    for sample in samples:
        if math.isclose(sample.route_distance_m, distance, abs_tol=1e-9):
            return sample
    return None


def _bracket(
    samples: Sequence[RouteWeatherSample], distance: float
) -> tuple[RouteWeatherSample, RouteWeatherSample]:
    for left, right in zip(samples, samples[1:]):
        if left.route_distance_m <= distance <= right.route_distance_m:
            return left, right
    raise ValueError("route distance is outside sample support")


def _sample_value(sample: RouteWeatherSample, variable: str) -> RouteWeatherValue:
    if sample.variables.get(variable) is None:
        return RouteWeatherValue(
            None,
            "UNKNOWN",
            "provider_route_sample",
            (sample.sample_id,),
            (0.0,),
            limitation_codes=(
                sample.missing_variable_reasons.get(variable, "WEATHER_PROVIDER_VALUE_MISSING"),
            ),
        )
    return RouteWeatherValue(
        sample.variables.get(variable),
        "PROVIDER_SAMPLE",
        "provider_route_sample",
        (sample.sample_id,),
        (0.0,),
        limitation_codes=("WEATHER_MODEL_GRID_NOT_ROUTE_POINT_PRECISION",),
    )


def _unknown(reason: str) -> RouteWeatherValue:
    return RouteWeatherValue(None, "UNKNOWN", "unsupported", (), (), limitation_codes=(reason,))


def _missing_variable_reasons(payload: dict[str, Any]) -> dict[str, str]:
    declared = dict(payload.get("missing_variable_reasons", {}))
    for variable in HOURLY_VARIABLES:
        if variable not in payload.get("variables", {}):
            declared.setdefault(variable, "WEATHER_VARIABLE_NOT_REQUESTED_BY_SOURCE_ARTIFACT")
    return declared


def _variable_semantics(payload: dict[str, Any]) -> dict[str, dict[str, Any]]:
    declared = dict(payload.get("variable_semantics", {}))
    for variable in payload.get("variables", {}):
        declared.setdefault(variable, weather_variable_semantics(variable))
    return {variable: semantics for variable, semantics in declared.items() if semantics}


def _circular_interpolate(first: float, second: float, fraction: float) -> float:
    radians_first, radians_second = math.radians(first), math.radians(second)
    x = (1 - fraction) * math.cos(radians_first) + fraction * math.cos(radians_second)
    y = (1 - fraction) * math.sin(radians_first) + fraction * math.sin(radians_second)
    value = math.degrees(math.atan2(y, x)) % 360
    return 0.0 if math.isclose(value, 360.0, abs_tol=1e-12) else value


def _angular_error(first: float, second: float) -> float:
    return abs((first - second + 180) % 360 - 180)


def _median(values: Sequence[float]) -> float | None:
    if not values:
        return None
    ordered = sorted(values)
    middle = len(ordered) // 2
    return ordered[middle] if len(ordered) % 2 else (ordered[middle - 1] + ordered[middle]) / 2


def _error_summary(values: Sequence[float]) -> dict[str, float | int | None]:
    if not values:
        return {
            "count": 0,
            "mae": None,
            "median_abs_error": None,
            "rmse": None,
            "max_abs_error": None,
        }
    return {
        "count": len(values),
        "mae": sum(values) / len(values),
        "median_abs_error": _median(values),
        "rmse": math.sqrt(sum(value * value for value in values) / len(values)),
        "max_abs_error": max(values),
    }


def _adjacent_mean(values: Sequence[float | int | None]) -> float | None:
    changes = [
        abs(float(right) - float(left))
        for left, right in zip(values, values[1:])
        if left is not None and right is not None
    ]
    return sum(changes) / len(changes) if changes else None
