"""Deterministic spatial sampling of a route at weather-model scale."""

from __future__ import annotations

from collections import defaultdict
from typing import Sequence

from mountain_twin.exposure import RoutePoint
from mountain_twin.route_analysis import prepare_route
from mountain_twin.weather.provider import WeatherLocation


def select_route_weather_locations(
    points: Sequence[RoutePoint], *, spacing_m: float = 2500
) -> tuple[WeatherLocation, ...]:
    """Choose route-scale samples, never one request per GPX point."""
    if spacing_m <= 0:
        raise ValueError("weather sample spacing must be positive")
    prepared = prepare_route(points).points
    if not prepared:
        raise ValueError("weather selection requires route points")
    selected: dict[int, set[str]] = defaultdict(set)
    total = prepared[-1].cumulative_distance_m
    targets = list(range(0, int(total), int(spacing_m))) + [int(total)]
    for target in targets:
        index = min(
            range(len(prepared)),
            key=lambda item: abs(prepared[item].cumulative_distance_m - target),
        )
        selected[index].add("regular_route_spacing_2_5km")
    selected[0].add("route_start")
    selected[len(prepared) - 1].add("route_end")
    elevations = [point.gpx_elevation_m for point in prepared]
    if all(value is not None for value in elevations):
        selected[elevations.index(min(elevations))].add("lowest_route_elevation")
        selected[elevations.index(max(elevations))].add("highest_route_elevation")
        ascent, descent = _one_km_transitions(prepared)
        selected[ascent].add("largest_one_km_ascent_transition")
        selected[descent].add("largest_one_km_descent_transition")
    locations = []
    for ordinal, index in enumerate(sorted(selected)):
        point = prepared[index]
        locations.append(
            WeatherLocation(
                sample_id=f"tmb_day_01_weather_{ordinal:02d}",
                point_index=point.point_index,
                route_distance_m=point.cumulative_distance_m,
                latitude=point.latitude,
                longitude=point.longitude,
                route_elevation_m=point.gpx_elevation_m,
                selection_reasons=tuple(sorted(selected[index])),
            )
        )
    return tuple(locations)


def _one_km_transitions(prepared) -> tuple[int, int]:
    deltas = []
    for index, point in enumerate(prepared):
        candidates = [
            previous
            for previous in range(index)
            if point.cumulative_distance_m - prepared[previous].cumulative_distance_m >= 1000
        ]
        if candidates:
            previous = candidates[-1]
            deltas.append((point.gpx_elevation_m - prepared[previous].gpx_elevation_m, index))
    if not deltas:
        return 0, len(prepared) - 1
    return max(deltas)[1], min(deltas)[1]
