"""Geometric propagation baseline: free-space loss + Bullington diffraction.

Model (per pilot -> drone path):
    PL = FSPL(slant distance, f) + J(nu)
where nu is the Fresnel-Kirchhoff parameter of the Bullington equivalent
knife-edge (ITU-R P.526, section 4.5, classic form without the smooth-earth
"delta" correction). J(nu) is the ITU-R P.526 knife-edge approximation:
    J(nu) = 6.9 + 20 log10( sqrt((nu-0.1)^2 + 1) + nu - 0.1 ),  nu > -0.78
    J(nu) = 0                                                   otherwise
nu <= -0.849 means >= 60 % first-Fresnel-zone clearance.

What is NOT modelled (by design, this is the physics prior):
  * vegetation loss, building penetration, multipath/reflections
  * interference / noise floor (handled by the link budget's rated-range levels)
  * antenna patterns and body blocking at the pilot

Implementation: radial rays from the pilot, profile sampled every `res` metres.
Earth curvature is applied with an effective radius k*a (k = 4/3). Heights are
expressed in the pilot's tangent frame (z - x^2/(2ka)); this is exactly
equivalent to the usual symmetric bulge x(d-x)/(2ka) because the difference
is linear in x and cancels against the Tx-Rx line.
"""
from __future__ import annotations

import math

import numba as nb
import numpy as np

C = 299_792_458.0
K_FACTOR = 4.0 / 3.0
A_EARTH = 6_371_000.0

NU_60_FRESNEL = -0.6 * math.sqrt(2.0)  # -0.849


def fspl_db(d_m, f_hz):
    """Free-space path loss (dB)."""
    d = np.maximum(np.asarray(d_m, dtype=np.float64), 1.0)
    return 20.0 * np.log10(d) + 20.0 * np.log10(f_hz) - 147.55


def knife_edge_db(nu):
    """ITU-R P.526 J(nu) approximation (dB)."""
    nu = np.asarray(nu, dtype=np.float64)
    v = np.maximum(nu, -0.78)  # J is 0 below -0.78; avoids log of ~0 for huge clearances
    out = 6.9 + 20.0 * np.log10(np.sqrt((v - 0.1) ** 2 + 1.0) + v - 0.1)
    return np.where(nu > -0.78, out, 0.0)


@nb.njit(cache=True, inline="always")
def _bilinear(a, r, c):
    n0, n1 = a.shape
    if r < 0.0:
        r = 0.0
    if c < 0.0:
        c = 0.0
    if r > n0 - 1.001:
        r = n0 - 1.001
    if c > n1 - 1.001:
        c = n1 - 1.001
    r0 = int(r)
    c0 = int(c)
    fr = r - r0
    fc = c - c0
    return ((1 - fr) * (1 - fc) * a[r0, c0] + (1 - fr) * fc * a[r0, c0 + 1]
            + fr * (1 - fc) * a[r0 + 1, c0] + fr * fc * a[r0 + 1, c0 + 1])


@nb.njit(parallel=True, cache=True)
def _radial_kernel(surface, terrain, center, res, n_rays, n_s, h_pilot,
                   alt, alt_agl, ke_radius):
    """Return (geom_nu, slant, collide) arrays of shape (n_rays, n_s).

    geom_nu is nu * sqrt(lambda): multiply by 1/sqrt(lambda) for any frequency.
    Sample index j corresponds to distance (j+1)*res.
    """
    geom = np.full((n_rays, n_s), np.nan)
    slant = np.full((n_rays, n_s), np.nan)
    collide = np.zeros((n_rays, n_s), dtype=np.bool_)
    zc = terrain[center, center]
    hts = zc + h_pilot

    for a in nb.prange(n_rays):
        th = 2.0 * math.pi * a / n_rays
        s, co = math.sin(th), math.cos(th)
        x = np.empty(n_s)
        Z = np.empty(n_s)       # surface in pilot tangent frame
        Tg = np.empty(n_s)      # bare terrain (no curvature) for AGL targets
        Sraw = np.empty(n_s)
        for i in range(n_s):
            d = (i + 1) * res
            col = center + d * s / res          # east
            row = center - d * co / res         # north is up (row decreases)
            x[i] = d
            Sraw[i] = _bilinear(surface, row, col)
            Tg[i] = _bilinear(terrain, row, col)
            Z[i] = Sraw[i] - d * d / (2.0 * ke_radius)

        stim = -1e30  # running max over i < j of (Z_i - hts)/x_i
        for j in range(n_s):
            d = x[j]
            base = Tg[j] if alt_agl else zc
            hr_abs = base + alt
            if hr_abs < Sraw[j] + 0.5:
                collide[a, j] = True
            hrs = hr_abs - d * d / (2.0 * ke_radius)
            slant[a, j] = math.sqrt(d * d + (hr_abs - hts) ** 2)

            if j == 0:
                geom[a, j] = -1e9  # no intermediate obstacle
            else:
                str_ = (hrs - hts) / d
                if stim < str_:
                    # line of sight: worst (largest) nu over the profile
                    best = -1e30
                    for i in range(j):
                        xi = x[i]
                        hline = (hts * (d - xi) + hrs * xi) / d
                        g = (Z[i] - hline) * math.sqrt(2.0 * d / (xi * (d - xi)))
                        if g > best:
                            best = g
                    geom[a, j] = best
                else:
                    srim = -1e30
                    for i in range(j):
                        v = (Z[i] - hrs) / (d - x[i])
                        if v > srim:
                            srim = v
                    xb = (hrs - hts + srim * d) / (stim + srim)
                    if xb <= 0.0:
                        xb = 1e-3
                    if xb >= d:
                        xb = d - 1e-3
                    hb = hts + stim * xb
                    hline = (hts * (d - xb) + hrs * xb) / d
                    geom[a, j] = (hb - hline) * math.sqrt(2.0 * d / (xb * (d - xb)))
            # update running max with point j for the next targets
            v = (Z[j] - hts) / x[j]
            if v > stim:
                stim = v
    return geom, slant, collide


def polar_to_grid(values, n, res, n_rays):
    """Resample (n_rays, n_s) polar arrays onto the n x n cartesian grid."""
    c = n // 2
    idx = (np.arange(n) - c) * res
    X, Y = np.meshgrid(idx, -idx)
    r = np.hypot(X, Y)
    th = np.mod(np.arctan2(X, Y), 2 * np.pi)  # clockwise from north
    ai = np.rint(th / (2 * np.pi) * n_rays).astype(int) % n_rays
    si = np.rint(r / res).astype(int) - 1
    n_s = values.shape[-1]
    valid = (si >= 0) & (si < n_s)
    out = np.full(values.shape[:-2] + (n, n), np.nan, dtype=values.dtype)
    out[..., valid] = values[..., ai[valid], si[valid]]
    return out


def path_loss_map(surface, terrain, res, alt_m, freqs_hz, h_pilot=1.2,
                  alt_agl=False, radius_m=None, k_factor=K_FACTOR):
    """Compute path-loss grids for one drone altitude.

    alt_agl=False: altitude is relative to the take-off point (DJI convention).
    Returns dict with 'pl' (n_freq, n, n) dB, 'nu' (n_freq, n, n),
    'collide' (n, n) bool, 'slant' (n, n) m.
    """
    n = surface.shape[0]
    center = n // 2
    if radius_m is None:
        radius_m = center * res
    n_s = int(radius_m // res) + 1
    n_rays = int(math.ceil(2 * math.pi * radius_m / res))
    geom, slant, collide = _radial_kernel(
        surface.astype(np.float64), terrain.astype(np.float64), center, float(res),
        n_rays, n_s, float(h_pilot), float(alt_m), bool(alt_agl), k_factor * A_EARTH)

    freqs = np.atleast_1d(np.asarray(freqs_hz, dtype=np.float64))
    lam = C / freqs
    nu = geom[None] / np.sqrt(lam)[:, None, None]
    pl = fspl_db(slant[None], freqs[:, None, None]) + knife_edge_db(nu)

    pl_g = polar_to_grid(pl, n, res, n_rays)
    nu_g = polar_to_grid(nu, n, res, n_rays)
    col_g = polar_to_grid(collide[None].astype(np.float32), n, res, n_rays)[0] > 0.5
    sl_g = polar_to_grid(slant[None], n, res, n_rays)[0]
    # the pilot's own cell
    pl_g[:, center, center] = np.nan
    return {"pl": pl_g, "nu": nu_g, "collide": col_g, "slant": sl_g}
