"""FastAPI server: pin a point, pick a device and altitude, get a coverage overlay."""
from __future__ import annotations

import time
from pathlib import Path
from typing import Literal

import numpy as np
from fastapi import FastAPI, HTTPException
from fastapi.responses import FileResponse
from pydantic import BaseModel, Field

from rcm.devices import DEVICES, INTERFERENCE_LEVELS
from rcm.service import (DEFAULT_ALTS, build_scene, coverage, min_altitude, overlay_bounds,
                         render_fresnel, render_margin, render_min_alt)

STATIC = Path(__file__).parent / "static"
app = FastAPI(title="Radio Coverage Map", version="0.1.0")


class CoverageRequest(BaseModel):
    lat: float = Field(ge=-85, le=85)
    lon: float = Field(ge=-180, le=180)
    device: str = "dji-mini-3-pro"
    mode: Literal["margin", "fresnel", "min_alt"] = "margin"
    alt_m: float = Field(60, ge=2, le=500)
    alt_agl: bool = False                 # False: above take-off (DJI convention)
    level: Literal["none", "low", "medium", "strong"] = "medium"
    h_pilot: float = Field(1.2, ge=0.3, le=100)  # controller height above ground
    radius_m: float = Field(2000, ge=200, le=5000)
    res_m: float = Field(5, ge=2, le=30)
    margin_req_db: float = 6.0
    buildings: bool = True


@app.get("/")
def index():
    return FileResponse(STATIC / "index.html")


@app.get("/api/devices")
def devices():
    return {"devices": [{"id": d.id, "name": d.name, "bands_ghz": [b / 1e9 for b in d.bands_hz],
                         "rated_range_m": d.rated_range_m, "source": d.source}
                        for d in DEVICES.values()],
            "levels": INTERFERENCE_LEVELS}


@app.post("/api/coverage")
def api_coverage(req: CoverageRequest):
    if req.device not in DEVICES:
        raise HTTPException(404, f"unknown device {req.device}")
    dev = DEVICES[req.device]
    t0 = time.time()
    # round the pin to ~1 m so repeated clicks reuse the cached scene
    scene = build_scene(round(req.lat, 5), round(req.lon, 5), req.radius_m, req.res_m, req.buildings)
    t_scene = time.time() - t0
    t1 = time.time()
    stats: dict = {}
    if req.mode == "min_alt":
        ma = min_altitude(scene, dev, req.level, req.margin_req_db, DEFAULT_ALTS,
                          req.h_pilot, req.alt_agl)
        img = render_min_alt(ma)
        n = ma.shape[0]
        yy, xx = np.mgrid[:n, :n] - n // 2
        inside = np.hypot(xx, yy) <= n // 2
        v = ~np.isnan(ma) & inside
        stats["reachable_pct"] = round(float(v.sum() / inside.sum() * 100), 1)
        stats["median_min_alt_m"] = float(np.median(ma[v])) if v.any() else None
    else:
        cov = coverage(scene, dev, req.alt_m, req.level, req.h_pilot, req.alt_agl)
        m = cov["margin"]
        v = ~np.isnan(m)
        stats["ok_pct"] = round(float((m[v] >= req.margin_req_db).mean() * 100), 1) if v.any() else 0
        stats["los_pct"] = round(float((cov["nu"][v] <= 0).mean() * 100), 1) if v.any() else 0
        stats["collide_pct"] = round(float(cov["collide"].mean() * 100), 1)
        img = render_margin(m, cov["collide"]) if req.mode == "margin" else render_fresnel(cov["nu"], cov["collide"])
    c = scene.grid.half
    k = max(1, int(25 / req.res_m))
    return {
        "image": img,
        "bounds": overlay_bounds(scene),
        "stats": stats,
        "pilot": {"ground_amsl_m": round(float(scene.terrain[c, c]), 1),
                  "max_building_within_25m": round(float((scene.surface - scene.terrain)[c - k:c + k + 1, c - k:c + k + 1].max()), 1)},
        "scene": {"grid": scene.terrain.shape[0], "res_m": req.res_m, "buildings": scene.n_buildings},
        "timing_s": {"scene": round(t_scene, 2), "compute": round(time.time() - t1, 2)},
    }
