"""High-level API: scene building, coverage maps, PNG rendering."""
from __future__ import annotations

import base64
import io
import time
from dataclasses import dataclass
from functools import lru_cache

import numpy as np
from matplotlib import colormaps
from PIL import Image

from .buildings import height_grid
from .devices import DEVICES, Device
from .grid import LocalGrid
from .propagation import NU_60_FRESNEL, path_loss_map
from .terrain import elevation_grid

DEFAULT_ALTS = tuple(range(10, 130, 10))  # 10..120 m


@dataclass
class Scene:
    grid: LocalGrid
    terrain: np.ndarray   # bare earth AMSL
    surface: np.ndarray   # terrain + buildings
    n_buildings: int
    build_s: float


@lru_cache(maxsize=16)
def build_scene(lat: float, lon: float, radius_m: float, res_m: float,
                use_buildings: bool = True) -> Scene:
    t0 = time.time()
    grid = LocalGrid(lat, lon, radius_m, res_m)
    terrain = elevation_grid(grid)
    if use_buildings:
        bh, nb_ = height_grid(grid)
    else:
        bh, nb_ = np.zeros_like(terrain), 0
    # the pilot stands on the ground: clear buildings in the pilot's own cells
    c = grid.half
    bh[c - 1:c + 2, c - 1:c + 2] = 0.0
    return Scene(grid, terrain, terrain + bh, nb_, time.time() - t0)


def coverage(scene: Scene, device: Device, alt_m: float, level: str,
             h_pilot: float = 1.2, alt_agl: bool = False) -> dict:
    """Best-band link margin (dB) at one altitude, plus diagnostics."""
    r = path_loss_map(scene.surface, scene.terrain, scene.grid.res_m, alt_m,
                      device.bands_hz, h_pilot=h_pilot, alt_agl=alt_agl,
                      radius_m=scene.grid.radius_m)
    margin_b = device.l_max(level)[:, None, None] - r["pl"]
    best = np.nanargmax(np.where(np.isnan(margin_b), -np.inf, margin_b), axis=0)
    margin = np.take_along_axis(margin_b, best[None], 0)[0]
    nu_best = np.take_along_axis(r["nu"], best[None], 0)[0]
    margin[r["collide"]] = np.nan
    return {"margin": margin, "nu": nu_best, "collide": r["collide"],
            "band": best, "pl": r["pl"]}


def min_altitude(scene: Scene, device: Device, level: str, margin_req: float = 6.0,
                 alts=DEFAULT_ALTS, h_pilot: float = 1.2, alt_agl: bool = False):
    """Lowest altitude (from `alts`) with margin >= margin_req; NaN if none."""
    out = np.full(scene.terrain.shape, np.nan)
    for a in sorted(alts):
        cov = coverage(scene, device, a, level, h_pilot, alt_agl)
        ok = (cov["margin"] >= margin_req) & np.isnan(out)
        out[ok] = a
    return out


# ---------------------------------------------------------------- rendering

def _rgba_png(rgba: np.ndarray) -> str:
    buf = io.BytesIO()
    Image.fromarray(rgba, "RGBA").save(buf, format="PNG", optimize=True)
    return "data:image/png;base64," + base64.b64encode(buf.getvalue()).decode()


def render_margin(margin: np.ndarray, collide: np.ndarray, vmin=-20, vmax=30, alpha=170):
    cmap = colormaps["RdYlGn"]
    z = np.clip((margin - vmin) / (vmax - vmin), 0, 1)
    rgba = (cmap(np.nan_to_num(z)) * 255).astype(np.uint8)
    rgba[..., 3] = alpha
    rgba[np.isnan(margin)] = (0, 0, 0, 0)
    rgba[collide] = (60, 60, 60, 200)
    return _rgba_png(rgba)


def render_min_alt(min_alt: np.ndarray, alts=DEFAULT_ALTS, alpha=170):
    cmap = colormaps["viridis_r"]
    lo, hi = min(alts), max(alts)
    z = np.clip((min_alt - lo) / max(hi - lo, 1), 0, 1)
    rgba = (cmap(np.nan_to_num(z)) * 255).astype(np.uint8)
    rgba[..., 3] = alpha
    rgba[np.isnan(min_alt)] = (200, 30, 30, 150)  # unreachable at any altitude
    # outside the computed radius: transparent
    n = min_alt.shape[0]
    yy, xx = np.mgrid[:n, :n] - n // 2
    rgba[np.hypot(xx, yy) > n // 2] = (0, 0, 0, 0)
    return _rgba_png(rgba)


def render_fresnel(nu: np.ndarray, collide: np.ndarray, alpha=170):
    rgba = np.zeros(nu.shape + (4,), np.uint8)
    rgba[nu <= NU_60_FRESNEL] = (40, 160, 60, alpha)            # >=60 % F1 clear
    rgba[(nu > NU_60_FRESNEL) & (nu <= 0)] = (240, 200, 40, alpha)  # LOS, F1 obstructed
    rgba[nu > 0] = (210, 50, 40, alpha)                          # no line of sight
    rgba[np.isnan(nu)] = (0, 0, 0, 0)
    rgba[collide] = (60, 60, 60, 200)
    return _rgba_png(rgba)


def overlay_bounds(scene: Scene):
    lon_w, lat_s, lon_e, lat_n = scene.grid.bbox_ll()
    return [[lat_s, lon_w], [lat_n, lon_e]]


__all__ = ["DEVICES", "build_scene", "coverage", "min_altitude", "render_margin",
           "render_min_alt", "render_fresnel", "overlay_bounds", "DEFAULT_ALTS"]
