"""Finite-range geodesic rays with a local flat-Earth elevation-angle model."""

from __future__ import annotations

import math
import statistics
from dataclasses import dataclass

from pyproj import Geod

from mountain_twin.terrain.provider import ElevationSource

GEOD = Geod(ellps="WGS84")


@dataclass(frozen=True)
class Horizon:
    horizon_angle_deg: float | None
    horizon_status: str
    sampled_count: int
    nodata_count: int
    out_of_bounds_count: int


def horizon_profile(
    dem: ElevationSource,
    latitude: float,
    longitude: float,
    azimuth_deg: float,
    step_m: float = 30,
    max_distance_m: float = 5000,
    observer_height_m: float = 0,
) -> Horizon:
    """Maximum sampled geometric angle from DEM origin plus nonnegative height.

    Zero height preserves the original model. Missing samples invalidate the ray;
    no curvature or atmospheric correction is applied.
    """
    if not all(math.isfinite(v) for v in (azimuth_deg, step_m, max_distance_m, observer_height_m)):
        raise ValueError("ray parameters must be finite")
    if step_m <= 0 or max_distance_m <= 0 or step_m > max_distance_m:
        raise ValueError("require 0 < step <= max distance")
    if observer_height_m < 0:
        raise ValueError("observer height must be nonnegative")
    count = math.ceil(max_distance_m / step_m)
    if count > 100000:
        raise ValueError("ray exceeds 100000 samples")
    origin = dem.sample(latitude, longitude)
    if origin.status != "valid":
        return Horizon(None, "origin_" + origin.status, 0, 0, 0)
    angles = []
    nodata = outside = 0
    for index in range(1, count + 1):
        distance = min(index * step_m, max_distance_m)
        lon, lat, _ = GEOD.fwd(longitude, latitude, azimuth_deg, distance)
        sample = dem.sample(lat, lon)
        if sample.status == "nodata":
            nodata += 1
        elif sample.status == "out_of_bounds":
            outside += 1
        else:
            angles.append(
                math.degrees(
                    math.atan2(
                        sample.elevation_m - origin.elevation_m - observer_height_m, distance
                    )
                )
            )
    # Missing samples can hide a higher ridge: never label an incomplete ray direct sun.
    complete = not (nodata or outside)
    return Horizon(
        max(angles) if complete else None,
        "complete" if complete else "incomplete",
        count,
        nodata,
        outside,
    )


def classify(solar_elevation_deg: float, horizon: Horizon) -> str:
    if solar_elevation_deg <= 0:
        return "astronomical_night"
    if horizon.horizon_angle_deg is None or horizon.horizon_status != "complete":
        return "terrain_shadow_unknown"
    return "terrain_shadow" if horizon.horizon_angle_deg >= solar_elevation_deg else "direct_sun"


def compare_elevations(points, samples) -> dict:
    if len(points) != len(samples):
        raise ValueError("point/sample count mismatch")
    differences = [
        s.elevation_m - p.elevation_m
        for p, s in zip(points, samples)
        if s.status == "valid" and p.elevation_m is not None
    ]
    return dict(
        difference_definition="DEM minus GPX (metres)",
        total_points=len(points),
        paired_points=len(differences),
        mean_difference_m=statistics.mean(differences) if differences else None,
        median_difference_m=statistics.median(differences) if differences else None,
        rmse_m=math.sqrt(statistics.mean(d * d for d in differences)) if differences else None,
        max_absolute_difference_m=max(map(abs, differences)) if differences else None,
        nodata_points=sum(s.status == "nodata" for s in samples),
        out_of_bounds_points=sum(s.status == "out_of_bounds" for s in samples),
        missing_gpx_elevation_points=sum(p.elevation_m is None for p in points),
    )
