"""Deterministic planned-route timelines and cached weather time interpolation."""

from __future__ import annotations

import json
import math
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, Sequence
from zoneinfo import ZoneInfo

from mountain_twin.snow.policy import MODELLED_SNOW_DEPTH_NEAREST_CONTEXT_POLICY_V0_1
from mountain_twin.weather.provider import HOURLY_VARIABLES, weather_variable_semantics
from mountain_twin.weather.spatial import RouteWeatherSample, WeatherField, interpolate_value

TEMPORAL_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",
    }
)
TEMPORAL_NEAREST_VARIABLES = frozenset(
    {
        "precipitation",
        "rain",
        "snowfall",
        "snow_depth",
        "precipitation_probability",
        "weather_code",
    }
)


@dataclass(frozen=True)
class Pause:
    route_distance_m: float
    duration_minutes: float


@dataclass(frozen=True)
class PlanningScenario:
    name: str
    start_datetime: datetime
    moving_speed_mps: float
    pauses: tuple[Pause, ...] = ()
    pace_factor: float | None = None

    def __post_init__(self) -> None:
        if self.start_datetime.utcoffset() is None:
            raise ValueError("planning start must be timezone-aware")
        if self.moving_speed_mps <= 0:
            raise ValueError("moving speed must be positive")
        if self.pace_factor is not None and self.pace_factor <= 0:
            raise ValueError("pace factor must be positive")
        if any(pause.route_distance_m < 0 or pause.duration_minutes < 0 for pause in self.pauses):
            raise ValueError("pause distance and duration must be non-negative")


@dataclass(frozen=True)
class TimelinePoint:
    point_index: int
    route_distance_m: float
    elevation_m: float | None
    planned_arrival: datetime
    elapsed_route_seconds: float
    cumulative_moving_seconds: float | None = None
    cumulative_pause_seconds: float = 0.0
    pace_interval_id: str | None = None


@dataclass(frozen=True)
class TemporalWeatherSample:
    base: RouteWeatherSample
    timezone: str
    times: tuple[datetime, ...]
    variables_by_time: dict[str, tuple[float | int | None, ...]]
    # Provenance for any variable taken from a named Open-Meteo model other
    # than the request's default selection ("best_match"), e.g.
    # {"freezing_level_height": "icon_seamless"}; empty = all default.
    variable_models: dict[str, str] = field(default_factory=dict)


@dataclass(frozen=True)
class TemporalWeatherField:
    base: WeatherField
    samples: tuple[TemporalWeatherSample, ...]


@dataclass(frozen=True)
class TemporalValue:
    value: float | int | None
    state: str
    contributing_times: tuple[str, ...]
    method: str
    limitation_codes: tuple[str, ...] = ()


def planning_scenarios(start_datetime: datetime) -> tuple[PlanningScenario, ...]:
    """Synthetic planning cases; these never claim to describe actual hiking pace."""
    return (
        PlanningScenario("FAST", start_datetime, 1.5, (Pause(12000, 20),), 1.20),
        PlanningScenario(
            "NOMINAL", start_datetime, 1.25, (Pause(9000, 20), Pause(19000, 25)), 1.00
        ),
        PlanningScenario(
            "SLOW", start_datetime, 1.0, (Pause(7000, 30), Pause(15000, 30), Pause(23000, 30)), 0.80
        ),
    )


def assign_timeline(
    prepared_points: Sequence[Any], scenario: PlanningScenario, pace_result: Any | None = None
) -> tuple[TimelinePoint, ...]:
    if pace_result is not None:
        moving_seconds_by_point = pace_result.moving_seconds_by_point
        interval_by_point = pace_result.interval_by_point
        expected_indices = tuple(point.point_index for point in prepared_points)
        if expected_indices != tuple(sorted(moving_seconds_by_point)):
            raise ValueError("pace result does not align with prepared route points")
    else:
        moving_seconds_by_point = {}
        interval_by_point = {}
    result = []
    for point in prepared_points:
        pause_seconds = sum(
            pause.duration_minutes * 60
            for pause in scenario.pauses
            if pause.route_distance_m <= point.cumulative_distance_m
        )
        moving_seconds = (
            moving_seconds_by_point[point.point_index]
            if pace_result is not None
            else point.cumulative_distance_m / scenario.moving_speed_mps
        )
        elapsed = moving_seconds + pause_seconds
        planned_arrival = (
            scenario.start_datetime.astimezone(timezone.utc) + timedelta(seconds=elapsed)
        ).astimezone(scenario.start_datetime.tzinfo)
        result.append(
            TimelinePoint(
                point_index=point.point_index,
                route_distance_m=point.cumulative_distance_m,
                elevation_m=point.gpx_elevation_m,
                planned_arrival=planned_arrival,
                elapsed_route_seconds=elapsed,
                cumulative_moving_seconds=moving_seconds,
                cumulative_pause_seconds=pause_seconds,
                pace_interval_id=interval_by_point.get(point.point_index),
            )
        )
    return tuple(result)


def load_temporal_field(summary_path: Path, cache_path: Path) -> TemporalWeatherField:
    """Join committed compact sample metadata to its ignored cached hourly response."""
    document = json.loads(summary_path.read_text(encoding="utf-8"))
    base = WeatherField(
        source_type=document["source_type"],
        scenario_datetime=document["scenario_datetime"],
        timezone=document["timezone"],
        provider=document["provider"],
        model_selection=document["model_selection"],
        samples=tuple(_sample_from_document(item) for item in document["samples"]),
    )
    raw = json.loads(cache_path.read_text(encoding="utf-8"))
    responses = raw if isinstance(raw, list) else [raw]
    if len(responses) != len(base.samples):
        raise ValueError("cached weather response count does not match compact sample count")
    zone = ZoneInfo(base.timezone)
    temporal = []
    for sample, response in zip(base.samples, responses):
        times = tuple(
            _provider_time(value, zone) for value in response.get("hourly", {}).get("time", ())
        )
        utc_times = tuple(value.astimezone(timezone.utc) for value in times)
        if any(right <= left for left, right in zip(utc_times, utc_times[1:])):
            raise ValueError("cached weather times must increase in UTC")
        hourly = response.get("hourly", {})
        variables = {variable: tuple(hourly.get(variable, ())) for variable in HOURLY_VARIABLES}
        temporal.append(TemporalWeatherSample(sample, base.timezone, times, variables))
    if not temporal or not temporal[0].times:
        raise ValueError("cached weather response has no hourly coverage")
    return TemporalWeatherField(base, tuple(temporal))


def temporal_value(
    sample: TemporalWeatherSample, instant: datetime, variable: str
) -> TemporalValue:
    if instant.tzinfo is None or instant.utcoffset() is None:
        raise ValueError("temporal weather instant must be timezone-aware")
    local = instant.astimezone(ZoneInfo(sample.timezone))
    instant_utc = local.astimezone(timezone.utc)
    sample_utc = tuple(value.astimezone(timezone.utc) for value in sample.times)
    if instant_utc < sample_utc[0] or instant_utc > sample_utc[-1]:
        return TemporalValue(
            None, "UNKNOWN", (), "unsupported", ("WEATHER_TIME_EXTRAPOLATION_UNSUPPORTED",)
        )
    exact_index = next(
        (index for index, value in enumerate(sample_utc) if value == instant_utc), None
    )
    if exact_index is not None:
        value = _series_value(sample, variable, exact_index)
        return TemporalValue(
            value,
            "PROVIDER_TIME" if value is not None else "UNKNOWN",
            (sample.times[exact_index].isoformat(),),
            "provider_hour",
            ()
            if value is not None
            else (
                sample.base.missing_variable_reasons.get(
                    variable, "WEATHER_PROVIDER_VALUE_MISSING"
                ),
            ),
        )
    left_index = max(index for index, value in enumerate(sample_utc) if value < instant_utc)
    right_index = left_index + 1
    first, second = (
        _series_value(sample, variable, left_index),
        _series_value(sample, variable, right_index),
    )
    times = (sample.times[left_index].isoformat(), sample.times[right_index].isoformat())
    if first is None or second is None:
        missing_reasons = tuple(
            sorted(
                {
                    sample.base.missing_variable_reasons.get(
                        variable, "WEATHER_TEMPORAL_INPUT_MISSING"
                    )
                    for value in (first, second)
                    if value is None
                }
            )
        )
        return TemporalValue(None, "UNKNOWN", times, "unsupported", missing_reasons)
    if variable in TEMPORAL_NEAREST_VARIABLES:
        chosen = (
            left_index
            if (instant_utc - sample_utc[left_index]) <= (sample_utc[right_index] - instant_utc)
            else right_index
        )
        return TemporalValue(
            _series_value(sample, variable, chosen),
            "NEAREST_CONTEXT",
            (sample.times[chosen].isoformat(),),
            "nearest_provider_hour_context",
            ("WEATHER_HOURLY_ACCUMULATION_OR_CATEGORICAL_CONTEXT",),
        )
    fraction = (instant_utc - sample_utc[left_index]).total_seconds() / (
        sample_utc[right_index] - sample_utc[left_index]
    ).total_seconds()
    if variable == "wind_direction_10m":
        value = _circular(float(first), float(second), fraction)
        method = "circular_vector_provider_time"
    elif variable in TEMPORAL_SCALAR_VARIABLES:
        value = float(first) + fraction * (float(second) - float(first))
        method = "linear_provider_time"
    else:
        return TemporalValue(
            None, "UNKNOWN", times, "unsupported", ("WEATHER_TEMPORAL_METHOD_UNDEFINED",)
        )
    return TemporalValue(
        value,
        "TEMPORAL_INTERPOLATED",
        times,
        method,
        ("WEATHER_TIME_INTERPOLATION_IS_NOT_NEW_PROVIDER_DATA",),
    )


@dataclass(frozen=True)
class TemporalWindowValue:
    """MIN/MAX of a real hourly variable over [window_start, window_end) at
    one location -- e.g. "minimum temperature between 19:00 and 07:00 here".

    Added alongside temporal_value() (one instant) and
    mountain_twin.weather.engine._ranges() (min/max across route points at
    one instant); this is the third, previously-missing axis: min/max across
    *time* at one point. See docs/design_reference/camps_staging_spike_v0_1.md
    section 2.1.
    """

    variable: str
    reduction: str
    value: float | int | None
    state: str
    contributing_times: tuple[str, ...]
    reason_codes: tuple[str, ...] = ()


def temporal_window_value(
    sample: TemporalWeatherSample,
    window_start: datetime,
    window_end: datetime,
    variable: str,
    reduction: str,
) -> TemporalWindowValue:
    """Reduce a real hourly series to its MIN or MAX within a time window.

    Never interpolates or invents a hypothetical value at the window edges:
    only real provider hours whose timestamp falls inside
    [window_start, window_end] contribute. A window that only partially
    overlaps the provider's covered hours still reduces over what does
    overlap, but is reported PARTIAL rather than COMPLETE; a window with no
    overlap at all is UNKNOWN, never a fabricated number.
    """
    if reduction not in ("MIN", "MAX"):
        raise ValueError("reduction must be MIN or MAX")
    if window_start.tzinfo is None or window_start.utcoffset() is None:
        raise ValueError("window start must be timezone-aware")
    if window_end.tzinfo is None or window_end.utcoffset() is None:
        raise ValueError("window end must be timezone-aware")
    if window_end <= window_start:
        raise ValueError("window end must be after window start")
    if not sample.times:
        return TemporalWindowValue(
            variable, reduction, None, "UNKNOWN", (), ("WEATHER_TIME_SERIES_EMPTY",)
        )

    sample_utc = tuple(value.astimezone(timezone.utc) for value in sample.times)
    start_utc = window_start.astimezone(timezone.utc)
    end_utc = window_end.astimezone(timezone.utc)

    covered_indices = [
        index for index, instant in enumerate(sample_utc) if start_utc <= instant <= end_utc
    ]
    if not covered_indices:
        return TemporalWindowValue(
            variable, reduction, None, "UNKNOWN", (), ("WEATHER_WINDOW_OUTSIDE_PROVIDER_COVERAGE",)
        )

    partial = start_utc < sample_utc[0] or end_utc > sample_utc[-1]
    reason_codes = ("WEATHER_WINDOW_PARTIALLY_OUTSIDE_PROVIDER_COVERAGE",) if partial else ()

    contributing_times: list[str] = []
    values: list[float | int] = []
    for index in covered_indices:
        value = _series_value(sample, variable, index)
        if value is None:
            continue
        values.append(value)
        contributing_times.append(sample.times[index].isoformat())

    if not values:
        missing_reason = sample.base.missing_variable_reasons.get(
            variable, "WEATHER_PROVIDER_VALUE_MISSING"
        )
        return TemporalWindowValue(
            variable, reduction, None, "UNKNOWN", (), reason_codes + (missing_reason,)
        )

    reduced = min(values) if reduction == "MIN" else max(values)
    state = "PARTIAL" if partial else "COMPLETE"
    return TemporalWindowValue(
        variable, reduction, reduced, state, tuple(contributing_times), reason_codes
    )


def combined_route_representation(
    field: TemporalWeatherField,
    timeline: Sequence[TimelinePoint],
    *,
    spatial_method: str = "policy",
) -> list[dict[str, Any]]:
    rows = []
    for point in timeline:
        temporally_sampled = []
        temporal_by_id = {}
        for sample in field.samples:
            values = {}
            for variable in HOURLY_VARIABLES:
                sampled = temporal_value(sample, point.planned_arrival, variable)
                values[variable] = sampled.value
                temporal_by_id[(sample.base.sample_id, variable)] = sampled
            temporally_sampled.append(
                RouteWeatherSample(
                    sample_id=sample.base.sample_id,
                    point_index=sample.base.point_index,
                    route_distance_m=sample.base.route_distance_m,
                    latitude=sample.base.latitude,
                    longitude=sample.base.longitude,
                    route_elevation_m=sample.base.route_elevation_m,
                    provider_elevation_m=sample.base.provider_elevation_m,
                    elevation_difference_m=sample.base.elevation_difference_m,
                    selection_reasons=sample.base.selection_reasons,
                    variables=values,
                    units=sample.base.units,
                    missing_variable_reasons=sample.base.missing_variable_reasons,
                    variable_semantics=sample.base.variable_semantics,
                )
            )
        values = {}
        for variable in HOURLY_VARIABLES:
            spatial = interpolate_value(
                temporally_sampled, point.route_distance_m, variable, method=spatial_method
            )
            temporal_records = [
                temporal_by_id[(sample_id, variable)]
                for sample_id in spatial.contributing_sample_ids
            ]
            temporal_state = (
                "UNKNOWN"
                if spatial.value is None
                or any(item.state == "UNKNOWN" for item in temporal_records)
                else "TEMPORAL_INTERPOLATED"
                if any(item.state == "TEMPORAL_INTERPOLATED" for item in temporal_records)
                else "NEAREST_CONTEXT"
                if any(item.state == "NEAREST_CONTEXT" for item in temporal_records)
                else "PROVIDER_TIME"
            )
            times = tuple(
                sorted({time for item in temporal_records for time in item.contributing_times})
            )
            values[variable] = {
                "value": spatial.value,
                "representation_state": "UNKNOWN"
                if spatial.value is None
                else spatial.representation_state,
                "spatial_method": spatial.method,
                "spatial_contributing_sample_ids": spatial.contributing_sample_ids,
                "spatial_support_distances_m": spatial.support_distances_m,
                "temporal_representation_state": temporal_state,
                "temporal_contributing_valid_times": times,
                "temporal_interpolation_method": ",".join(
                    sorted({item.method for item in temporal_records})
                )
                if temporal_records
                else "unsupported",
                "field_semantics": _field_semantics(
                    field.samples, spatial.contributing_sample_ids, variable
                ),
                "elevation_adjustment_state": "not_applied",
                "limitation_codes": tuple(
                    sorted(
                        set(spatial.limitation_codes)
                        | {code for item in temporal_records for code in item.limitation_codes}
                    )
                ),
                **_snow_projection_metadata(variable, spatial, temporal_records),
            }
        rows.append(
            {
                "point_index": point.point_index,
                "route_distance_m": point.route_distance_m,
                "elevation_m": point.elevation_m,
                "planned_arrival": point.planned_arrival.isoformat(),
                "elapsed_route_seconds": point.elapsed_route_seconds,
                "scenario_datetime_timezone": field.base.timezone,
                "provider": field.base.provider,
                "source_type": field.base.source_type,
                "values": values,
            }
        )
    return rows


def _sample_from_document(item: dict[str, Any]) -> RouteWeatherSample:
    selection, payload = item["selection"], item["contract"]["payload"]
    return 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
                for variable, semantics in payload.get("variable_semantics", {}).items()
                if semantics
            },
            **{
                variable: weather_variable_semantics(variable)
                for variable in payload.get("variables", {})
                if variable not in payload.get("variable_semantics", {})
                and weather_variable_semantics(variable)
            },
        },
    )


def _provider_time(value: str, zone: ZoneInfo) -> datetime:
    parsed = datetime.fromisoformat(value)
    return parsed.replace(tzinfo=zone) if parsed.tzinfo is None else parsed.astimezone(zone)


def _series_value(sample: TemporalWeatherSample, variable: str, index: int) -> float | int | None:
    series = sample.variables_by_time.get(variable, ())
    return series[index] if index < len(series) else None


def _snow_projection_metadata(
    variable: str, spatial: Any, temporal_records: Sequence[TemporalValue]
) -> dict[str, Any]:
    if variable != "snow_depth":
        return {}
    temporal = temporal_records[0] if len(temporal_records) == 1 else None
    return {
        "projection_policy_id": MODELLED_SNOW_DEPTH_NEAREST_CONTEXT_POLICY_V0_1.policy_id,
        "snow_spatial_representation": (
            "EXACT_SOURCE_SUPPORT"
            if spatial.representation_state == "PROVIDER_SAMPLE"
            else "NEAREST_PROVIDER_ROUTE_SAMPLE"
            if spatial.representation_state == "NEAREST_CONTEXT"
            else "UNKNOWN"
        ),
        "snow_temporal_projection": (
            "EXACT_SOURCE_SUPPORT"
            if temporal is not None and temporal.state == "PROVIDER_TIME"
            else "NEAREST_PROVIDER_HOUR"
            if temporal is not None and temporal.state == "NEAREST_CONTEXT"
            else "UNKNOWN"
        ),
        "coverage": "FULL" if spatial.value is not None else "UNAVAILABLE",
    }


def _field_semantics(
    samples: Sequence[TemporalWeatherSample], contributing_sample_ids: Sequence[str], variable: str
) -> dict[str, Any] | None:
    selected = [
        sample.base.variable_semantics.get(variable)
        for sample in samples
        if sample.base.sample_id in contributing_sample_ids
        and sample.base.variable_semantics.get(variable)
    ]
    return selected[0] if selected and all(item == selected[0] for item in selected) else None


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 _circular(first: float, second: float, fraction: float) -> float:
    x = (1 - fraction) * math.cos(math.radians(first)) + fraction * math.cos(math.radians(second))
    y = (1 - fraction) * math.sin(math.radians(first)) + fraction * math.sin(math.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
