"""Copernicus DEM GLO-90 tiles from the public AWS open-data bucket, cached on
disk, read as one regular grid over a bounding box.

docs/design_reference/terrain3d_sun_spike_v0_1.md section c: the terrain
behind the 3D view's cast shadow. No key or account; one 1 x 1 degree
Cloud-Optimized GeoTIFF per tile (about 5 MB), downloaded once into the
cache directory -- after that the grid is read fully offline.

GLO-90 is a surface model (DSM: forest canopy and buildings included), not
bare earth; heights are EGM2008 orthometric metres. A tile the bucket does
not have is open sea in this dataset (no land, elevation 0 m); any other
failure to get a tile is a failure of the whole grid, never a guessed height.
"""

from __future__ import annotations

import math
from dataclasses import dataclass
from pathlib import Path
from urllib.error import HTTPError

import numpy as np
import rasterio

from mountain_twin.offline import urlopen  # AV-045: offline mode for tests

GLO90_BUCKET_URL = "https://copernicus-dem-90m.s3.amazonaws.com"
GLO90_PRODUCT = "COP-DEM_GLO-90-DGED (Copernicus DEM, 3 arc seconds)"
# The dataset's own posting: 3 arc seconds of latitude everywhere.
GLO90_ARC_DEGREES = 1.0 / 1200.0


class TerrainGridUnavailable(RuntimeError):
    """The elevation grid could not be obtained (download or read failed)."""


@dataclass(frozen=True)
class TerrainGrid:
    """Elevations [m] on a regular latitude/longitude grid, row 0 at the
    north edge; (west, north) is the outer corner of cell [0, 0]."""

    elevation_m: np.ndarray
    west: float
    north: float
    cell_degrees: float

    @property
    def bounds(self) -> tuple[float, float, float, float]:
        rows, cols = self.elevation_m.shape
        return (
            self.west,
            self.north - rows * self.cell_degrees,
            self.west + cols * self.cell_degrees,
            self.north,
        )


def glo90_tile_name(latitude_floor: int, longitude_floor: int) -> str:
    north = "N" if latitude_floor >= 0 else "S"
    east = "E" if longitude_floor >= 0 else "W"
    return (
        f"Copernicus_DSM_COG_30_{north}{abs(latitude_floor):02d}_00_"
        f"{east}{abs(longitude_floor):03d}_00_DEM"
    )


class CopernicusGlo90:
    """``grid(west, south, east, north)`` over any area; tiles cached in
    ``cache_directory``. ``opener`` is injectable for tests."""

    source = "Copernicus DEM GLO-90 (AWS Open Data)"
    product = GLO90_PRODUCT
    surface_type = "DSM"
    nominal_cell_m = 90.0

    def __init__(self, cache_directory: Path, *, opener=urlopen) -> None:
        self.cache_directory = Path(cache_directory)
        self._opener = opener
        self.download_count = 0  # observability/test hook: real downloads only

    def grid(self, west: float, south: float, east: float, north: float) -> TerrainGrid:
        if not (west < east and south < north):
            raise ValueError("terrain grid bounds must have positive extent")
        cell = GLO90_ARC_DEGREES
        # Snap to the dataset's own posting so cells line up with its pixels
        # (the small tolerance keeps a bound that already lies on the
        # posting from gaining a cell through floating-point error).
        tolerance = 1e-7
        west_index = math.floor(west / cell + tolerance)
        north_index = math.ceil(north / cell - tolerance)
        west_edge, north_edge = west_index * cell, north_index * cell
        cols = max(2, math.ceil(east / cell - tolerance) - west_index)
        rows = max(2, north_index - math.floor(south / cell + tolerance))
        elevation = np.zeros((rows, cols), dtype=np.float32)
        longitudes = west_edge + (np.arange(cols) + 0.5) * cell
        latitudes = north_edge - (np.arange(rows) + 0.5) * cell
        for latitude_floor in range(math.floor(latitudes.min()), math.floor(latitudes.max()) + 1):
            for longitude_floor in range(
                math.floor(longitudes.min()), math.floor(longitudes.max()) + 1
            ):
                row_index = np.nonzero(np.floor(latitudes) == latitude_floor)[0]
                col_index = np.nonzero(np.floor(longitudes) == longitude_floor)[0]
                if not len(row_index) or not len(col_index):
                    continue
                path = self._tile(latitude_floor, longitude_floor)
                if path is None:  # no tile in the dataset: open sea
                    continue
                elevation[np.ix_(row_index, col_index)] = self._read(
                    path, latitudes[row_index], longitudes[col_index]
                )
        return TerrainGrid(elevation, west_edge, north_edge, cell)

    @staticmethod
    def _read(path: Path, latitudes: np.ndarray, longitudes: np.ndarray) -> np.ndarray:
        """Nearest tile pixel for each cell centre (tiles north of 50 degrees
        have fewer columns; the dataset's own transform says which)."""
        try:
            with rasterio.open(path) as dataset:
                inverse = ~dataset.transform
                tile_cols = np.clip(
                    np.floor(inverse.a * longitudes + inverse.c).astype(int), 0, dataset.width - 1
                )
                tile_rows = np.clip(
                    np.floor(inverse.e * latitudes + inverse.f).astype(int), 0, dataset.height - 1
                )
                window = rasterio.windows.Window(
                    int(tile_cols.min()),
                    int(tile_rows.min()),
                    int(tile_cols.max() - tile_cols.min() + 1),
                    int(tile_rows.max() - tile_rows.min() + 1),
                )
                data = dataset.read(1, window=window).astype(np.float32)
        except rasterio.errors.RasterioError as error:
            raise TerrainGridUnavailable("TERRAIN_DEM_READ_FAILED") from error
        values = data[np.ix_(tile_rows - tile_rows.min(), tile_cols - tile_cols.min())]
        return np.where(np.isfinite(values), values, 0.0)

    def _tile(self, latitude_floor: int, longitude_floor: int) -> Path | None:
        name = glo90_tile_name(latitude_floor, longitude_floor)
        path = self.cache_directory / f"{name}.tif"
        sea = self.cache_directory / f"{name}.absent"
        if path.exists():
            return path
        if sea.exists():
            return None
        self.download_count += 1
        try:
            with self._opener(f"{GLO90_BUCKET_URL}/{name}/{name}.tif", timeout=60) as response:
                body = response.read()
        except HTTPError as error:
            if error.code == 404:  # the bucket has no such tile: open sea
                self.cache_directory.mkdir(parents=True, exist_ok=True)
                sea.write_text("no tile in the dataset (open sea)\n", encoding="utf-8")
                return None
            raise TerrainGridUnavailable("TERRAIN_DEM_DOWNLOAD_FAILED") from error
        except OSError as error:
            raise TerrainGridUnavailable("TERRAIN_DEM_DOWNLOAD_FAILED") from error
        self.cache_directory.mkdir(parents=True, exist_ok=True)
        partial = path.with_suffix(".part")
        partial.write_bytes(body)
        partial.replace(path)
        return path
