"""The sun and the terrain's cast shadow around a route at one moment -- the
data behind the 3D view's light (docs/design_reference/
terrain3d_sun_spike_v0_1.md sections c-d).

Two facts, kept apart (ADR-003):

- the sun's position: a deterministic astronomical calculation
  (mountain_twin.solar.sun_engine.solar_position, no refraction);
- the cast shadow: derived from a terrain *model* (Copernicus GLO-90, a
  surface model) -- every grid cell from which the terrain towards the sun
  rises above the sun's ray, on the same local flat-Earth geometry as
  mountain_twin.terrain.horizon (no curvature, no refraction).

The shadow is a picture of the model, not an observation and not a verdict
about the route (ADR-001): it goes to the map as a raster, and nothing here
says whether a place on the route is "in shadow". The browser only draws
what it is given. A sun below the horizon casts no shadow (NOT_APPLICABLE);
terrain that cannot be read is UNAVAILABLE -- never a guessed picture.
"""

from __future__ import annotations

import base64
import math
import struct
import zlib
from datetime import datetime
from typing import Any, Sequence

import numpy as np

from mountain_twin.solar.states import classify_astronomical_state
from mountain_twin.solar.sun_engine import solar_position
from mountain_twin.terrain.copernicus import TerrainGrid, TerrainGridUnavailable

TERRAIN_LIGHT_CONTRACT = "route_terrain_light_v0_1"
# The terrain kept around the route: what can shade it, and the shadows the
# view shows next to it. A shadow cast from farther away is not seen.
SHADOW_MARGIN_M = 8000.0
# Long routes: the grid is thinned to stay answerable while a slider moves
# (the spike's 291 x 232 cells took 0.06-0.24 s per moment).
MAX_GRID_CELLS = 400_000
METRES_PER_DEGREE = 111_320.0

LIMITATIONS = (
    "TERRAIN_MODEL_IS_A_SURFACE_MODEL_NOT_BARE_EARTH",
    "LOCAL_FLAT_EARTH_NO_CURVATURE_NO_REFRACTION",
    "SHADOWS_CAST_FROM_BEYOND_THE_MARGIN_NOT_INCLUDED",
    "SUN_POSITION_AT_THE_ROUTE_CENTRE_FOR_THE_WHOLE_AREA",
)


def cast_shadow(
    elevation_m: np.ndarray,
    cell_east_m: float,
    cell_north_m: float,
    sun_azimuth_deg: float,
    sun_elevation_deg: float,
    max_reach_m: float,
) -> np.ndarray:
    """True for every cell whose ray towards the sun is blocked by terrain:
    marching from the cell along the sun's azimuth, some (bilinearly
    interpolated) elevation on the way is above the ray. Row 0 is north."""
    if sun_elevation_deg <= 0:
        raise ValueError("a cast shadow needs the sun above the horizon")
    rows, cols = elevation_m.shape
    step = min(cell_east_m, cell_north_m)
    east = math.sin(math.radians(sun_azimuth_deg))
    north = math.cos(math.radians(sun_azimuth_deg))
    rise_per_metre = math.tan(math.radians(sun_elevation_deg))
    relief = float(elevation_m.max() - elevation_m.min())
    # Beyond this distance even the highest terrain is under the ray.
    reach = min(max_reach_m, relief / rise_per_metre)
    shadow = np.zeros(elevation_m.shape, dtype=bool)
    row_index, col_index = np.mgrid[0:rows, 0:cols]
    for k in range(1, int(reach / step) + 2):
        distance = k * step
        x = col_index + east * distance / cell_east_m
        y = row_index - north * distance / cell_north_m
        x0, y0 = np.floor(x).astype(np.intp), np.floor(y).astype(np.intp)
        inside = (x0 >= 0) & (x0 < cols - 1) & (y0 >= 0) & (y0 < rows - 1)
        if not inside.any():
            break
        x0c, y0c = np.clip(x0, 0, cols - 2), np.clip(y0, 0, rows - 2)
        fx, fy = x - x0, y - y0
        terrain = (
            elevation_m[y0c, x0c] * (1 - fx) * (1 - fy)
            + elevation_m[y0c, x0c + 1] * fx * (1 - fy)
            + elevation_m[y0c + 1, x0c] * (1 - fx) * fy
            + elevation_m[y0c + 1, x0c + 1] * fx * fy
        )
        shadow |= inside & (terrain > elevation_m + distance * rise_per_metre)
    return shadow


def shadow_png(shadow: np.ndarray) -> bytes:
    """The mask as an RGBA PNG: opaque black where in shadow, transparent
    elsewhere (the map decides how dark a shadow looks)."""
    rows, cols = shadow.shape
    rgba = np.zeros((rows, cols, 4), dtype=np.uint8)
    rgba[..., 3] = np.where(shadow, 255, 0)
    raw = b"".join(b"\x00" + rgba[row].tobytes() for row in range(rows))

    def chunk(kind: bytes, data: bytes) -> bytes:
        return (
            struct.pack(">I", len(data))
            + kind
            + data
            + struct.pack(">I", zlib.crc32(kind + data) & 0xFFFFFFFF)
        )

    return (
        b"\x89PNG\r\n\x1a\n"
        + chunk(b"IHDR", struct.pack(">IIBBBBB", cols, rows, 8, 6, 0, 0, 0))
        + chunk(b"IDAT", zlib.compress(raw, 6))
        + chunk(b"IEND", b"")
    )


def route_bounds(route_points: Sequence[Any], margin_m: float) -> tuple[float, float, float, float]:
    """(west, south, east, north) of the route widened by ``margin_m``."""
    latitudes = [point.latitude for point in route_points]
    longitudes = [point.longitude for point in route_points]
    centre = math.radians((min(latitudes) + max(latitudes)) / 2)
    d_lat = margin_m / METRES_PER_DEGREE
    d_lon = margin_m / (METRES_PER_DEGREE * max(0.05, math.cos(centre)))
    return (
        min(longitudes) - d_lon,
        min(latitudes) - d_lat,
        max(longitudes) + d_lon,
        max(latitudes) + d_lat,
    )


def terrain_light_document(
    *,
    journey_id: str,
    route_points: Sequence[Any],
    moment: datetime,
    timezone_name: str,
    dem_provider: Any,
) -> dict[str, Any]:
    if not route_points:
        raise ValueError("terrain light needs route points")
    if moment.utcoffset() is None:
        raise ValueError("terrain light needs a timezone-aware moment")
    west, south, east, north = route_bounds(route_points, SHADOW_MARGIN_M)
    centre_lat, centre_lon = (south + north) / 2, (west + east) / 2
    sun_elevation, sun_azimuth = solar_position(moment, centre_lat, centre_lon)
    document: dict[str, Any] = {
        "contract": TERRAIN_LIGHT_CONTRACT,
        "journey_id": journey_id,
        "time": moment.isoformat(),
        "timezone": timezone_name,
        "sun": {
            "evidence_type": "DETERMINISTIC_CALCULATION",
            "method": "NOAA fractional-year solar position, geometric (no refraction)",
            "at": {"latitude": centre_lat, "longitude": centre_lon},
            "elevation_deg": round(sun_elevation, 2),
            "azimuth_deg": round(sun_azimuth, 2),
            "state": classify_astronomical_state(sun_elevation).value,
        },
        "terrain": {
            "source": dem_provider.source,
            "product": dem_provider.product,
            "surface_type": dem_provider.surface_type,
            "margin_m": SHADOW_MARGIN_M,
        },
        "limitations": list(LIMITATIONS),
    }
    if sun_elevation <= 0:
        document["shadow"] = {"state": "NOT_APPLICABLE", "reason": "SUN_BELOW_HORIZON"}
        return document
    try:
        grid: TerrainGrid = dem_provider.grid(west, south, east, north)
    except TerrainGridUnavailable as unavailable:
        document["shadow"] = {"state": "UNAVAILABLE", "reason": str(unavailable)}
        return document
    elevation = grid.elevation_m
    stride = max(1, math.ceil(math.sqrt(elevation.size / MAX_GRID_CELLS)))
    elevation = elevation[::stride, ::stride]
    cell_degrees = grid.cell_degrees * stride
    cell_north_m = cell_degrees * METRES_PER_DEGREE
    cell_east_m = cell_north_m * math.cos(math.radians(centre_lat))
    shadow = cast_shadow(
        elevation, cell_east_m, cell_north_m, sun_azimuth, sun_elevation, SHADOW_MARGIN_M
    )
    rows, cols = shadow.shape
    document["terrain"]["cell_m"] = round(cell_north_m, 1)
    document["shadow"] = {
        "state": "AVAILABLE",
        "evidence_type": "DERIVED_FROM_TERRAIN_MODEL",
        "method": "cast shadow: terrain above the sun's ray along its azimuth, bilinear",
        # (west, south, east, north) of the picture's outer edges.
        "bounds": [
            grid.west,
            grid.north - rows * cell_degrees,
            grid.west + cols * cell_degrees,
            grid.north,
        ],
        "rows": rows,
        "columns": cols,
        "shadow_fraction": round(float(shadow.mean()), 4),
        "image_media_type": "image/png",
        "image_base64": base64.b64encode(shadow_png(shadow)).decode("ascii"),
    }
    return document
