"""Elevations for "Pogoda" (AV-062) from the DEM the server already reads for
the 3D view: Copernicus GLO-90 (mountain_twin.terrain.copernicus; a surface
model, 90 m posting, EGM2008 metres).

* a place without OSM's ``ele``: the DEM's height at the point; a peak's: the
  highest DEM cell within SUMMIT_RADIUS_M (a peak's coordinates are often a
  few tens of metres off its top, and the 90 m cells smooth the top down);
* a peak's forecast heights, like Mountain-Forecast's: the summit, its base
  -- the BASE_PERCENTILE height of the terrain within BASE_RADIUS_M, i.e.
  the valley floors around it (for Śnieżka: Karpacz's side) -- and the
  middle between them, each rounded to LEVEL_ROUNDING_M. When the summit
  rises less than MIN_RELIEF_M above that base, one height only.

Every height says where it came from (``elevation_source``): OSM's tag, the
DEM at the point, the DEM's maximum, the DEM's percentile, or the midpoint
of two of those. No DEM (a download failed): the summit alone, with the
reason -- never a guessed base.
"""

from __future__ import annotations

import math
from dataclasses import dataclass
from typing import Any

import numpy as np

SUMMIT_RADIUS_M = 150.0
BASE_RADIUS_M = 6000.0
BASE_PERCENTILE = 10.0
MIN_RELIEF_M = 300.0
LEVEL_ROUNDING_M = 50.0
POINT_WINDOW_DEGREES = 0.003
METRES_PER_DEGREE = 111_320.0

SOURCE_OSM = "OSM_ELE"
SOURCE_DEM_POINT = "DEM_GLO90_POINT"
SOURCE_DEM_SUMMIT = "DEM_GLO90_MAX_150M"
SOURCE_DEM_BASE = "DEM_GLO90_P10_6KM"
SOURCE_MIDPOINT = "MIDPOINT_SUMMIT_BASE"
SOURCE_PROVIDER = "PROVIDER_GRID"


@dataclass(frozen=True)
class Level:
    level_id: str
    label: str
    elevation_m: float | None
    elevation_source: str

    def to_dict(self) -> dict[str, Any]:
        return {
            "level_id": self.level_id,
            "label": self.label,
            "elevation_m": self.elevation_m,
            "elevation_source": self.elevation_source,
        }


def _window(latitude: float, longitude: float, radius_m: float) -> tuple[float, ...]:
    dlat = radius_m / METRES_PER_DEGREE
    dlon = radius_m / (METRES_PER_DEGREE * max(0.05, math.cos(math.radians(latitude))))
    return longitude - dlon, latitude - dlat, longitude + dlon, latitude + dlat


def _cells_within(grid, latitude: float, longitude: float, radius_m: float) -> np.ndarray:
    """The grid's heights whose cell centres lie within ``radius_m``."""
    rows, cols = grid.elevation_m.shape
    lats = grid.north - (np.arange(rows) + 0.5) * grid.cell_degrees
    lons = grid.west + (np.arange(cols) + 0.5) * grid.cell_degrees
    dy = (lats[:, None] - latitude) * METRES_PER_DEGREE
    dx = (lons[None, :] - longitude) * METRES_PER_DEGREE * math.cos(math.radians(latitude))
    inside = dx * dx + dy * dy <= radius_m * radius_m
    return grid.elevation_m[inside]


def dem_point(dem, latitude: float, longitude: float) -> float | None:
    """The DEM's height of the cell holding the point (None: no DEM)."""
    try:
        grid = dem.grid(*_window(latitude, longitude, POINT_WINDOW_DEGREES * METRES_PER_DEGREE))
    except Exception:  # noqa: BLE001 -- any DEM failure is "no DEM height"
        return None
    rows, cols = grid.elevation_m.shape
    row = min(rows - 1, max(0, int((grid.north - latitude) / grid.cell_degrees)))
    col = min(cols - 1, max(0, int((longitude - grid.west) / grid.cell_degrees)))
    return float(round(float(grid.elevation_m[row, col])))


def dem_summit(dem, latitude: float, longitude: float) -> float | None:
    try:
        grid = dem.grid(*_window(latitude, longitude, SUMMIT_RADIUS_M * 1.5))
    except Exception:  # noqa: BLE001
        return None
    cells = _cells_within(grid, latitude, longitude, SUMMIT_RADIUS_M)
    return float(round(float(cells.max()))) if cells.size else None


def dem_base(dem, latitude: float, longitude: float) -> float | None:
    try:
        grid = dem.grid(*_window(latitude, longitude, BASE_RADIUS_M))
    except Exception:  # noqa: BLE001
        return None
    cells = _cells_within(grid, latitude, longitude, BASE_RADIUS_M)
    return float(np.percentile(cells, BASE_PERCENTILE)) if cells.size else None


def _rounded(value: float) -> float:
    return float(round(value / LEVEL_ROUNDING_M) * LEVEL_ROUNDING_M)


def place_elevation(
    kind: str, latitude: float, longitude: float, osm_ele: float | None, dem
) -> tuple[float | None, str]:
    """The place's own height and where it is from."""
    if osm_ele is not None:
        return osm_ele, SOURCE_OSM
    value = (dem_summit if kind == "PEAK" else dem_point)(dem, latitude, longitude)
    if value is not None:
        return value, SOURCE_DEM_SUMMIT if kind == "PEAK" else SOURCE_DEM_POINT
    # No height known: the provider answers for its own grid cell's height.
    return None, SOURCE_PROVIDER


def forecast_levels(
    kind: str,
    latitude: float,
    longitude: float,
    elevation_m: float | None,
    elevation_source: str,
    dem,
) -> tuple[list[Level], str | None]:
    """The heights the forecast is asked for, highest first, and why there
    is only one when a peak gets one (None otherwise)."""
    if kind != "PEAK":
        return [Level("place", "Miejsce", elevation_m, elevation_source)], None
    summit = Level("summit", "Szczyt", elevation_m, elevation_source)
    if elevation_m is None:
        return [summit], "SUMMIT_ELEVATION_UNAVAILABLE"
    base = dem_base(dem, latitude, longitude)
    if base is None:
        return [summit], "BASE_ELEVATION_UNAVAILABLE"
    base = _rounded(base)
    if elevation_m - base < MIN_RELIEF_M:
        return [summit], "RELIEF_BELOW_THRESHOLD"
    middle = _rounded((elevation_m + base) / 2)
    return [
        summit,
        Level("middle", "Środek", middle, SOURCE_MIDPOINT),
        Level("base", "Podnóże", base, SOURCE_DEM_BASE),
    ], None
