"""RadioMapSeer loading + cheap physics features.

Encoding (RadioUNet paper, section 3.3): gray f in [0,1], PL_dB = -147 + 99.16 * f.
Buildings png: 255 = building. Antenna png: single 255 pixel at the Tx.

Paper metric: the official RadioUNet loader clips targets at the noise floor THRESH = 0.2 gray
(-127.2 dB) and rescales (f - 0.2) / 0.8, so one "paper gray" unit is 0.8 * 99.16 = 79.3 dB
(the paper's "dB = 80 x gray"). See docs/research/03-mps-runs.md.
"""
from __future__ import annotations

import os
from pathlib import Path

import numba as nb
import numpy as np
from PIL import Image

ROOT = Path(os.environ.get("RMS_ROOT", Path(__file__).resolve().parent.parent / "data/rms"))
P_TRNC, M1 = -147.0, -47.84
DB_RANGE = M1 - P_TRNC  # 99.16
N = 256
THRESH = 0.2  # RadioUNet noise-floor threshold in gray (P_thr = -147 + 0.2 * 99.16 = -127.2 dB)
PAPER_DB_PER_GRAY = (1 - THRESH) * DB_RANGE  # 79.33


def gray_to_db(f):
    return P_TRNC + DB_RANGE * np.asarray(f, dtype=np.float32)


def load_map(m: int) -> np.ndarray:
    return (np.asarray(Image.open(ROOT / f"png/buildings_complete/{m}.png")) > 0).astype(np.float32)


def load_tx(m: int, t: int) -> tuple[int, int]:
    a = np.asarray(Image.open(ROOT / f"png/antennas/{m}_{t}.png"))
    r, c = np.argwhere(a > 0)[0]
    return int(r), int(c)


def load_gain(m: int, t: int, sim: str = "DPM") -> np.ndarray:
    return np.asarray(Image.open(ROOT / f"gain/{sim}/{m}_{t}.png"), dtype=np.float32) / 255.0


def split(seed: int = 0, official: bool = False):
    """Map ids for train / val / test.

    official=True: RadioUNet's own split (github.com/RonLevie/RadioUNet, lib/loaders.py): shuffle
    0..699 with np.random.seed(42), add 1 (so map 0 is the one left out, not 700), then take the
    inclusive index ranges 0-500 / 501-600 / 601-699, i.e. 501 / 100 / 99 maps.
    official=False: our older seeded 500/100/100 split of ids 0..699.
    """
    if official:
        ids = np.arange(0, 700, 1, dtype=np.int16)
        np.random.RandomState(42).shuffle(ids)  # same stream as np.random.seed(42); np.random.shuffle
        ids = ids.astype(np.int64) + 1
        return ids[0:501], ids[501:601], ids[601:700]
    ids = np.random.default_rng(seed).permutation(700)
    return ids[:500], ids[500:600], ids[600:700]


@nb.njit(cache=True)
def _line_features(bld, tr, tc):
    """For every pixel: straight-line metres inside buildings and number of wall crossings from the Tx."""
    n = bld.shape[0]
    inside = np.zeros((n, n), np.float32)
    walls = np.zeros((n, n), np.float32)
    for r in range(n):
        for c in range(n):
            dr, dc = r - tr, c - tc
            L = max(abs(dr), abs(dc))
            if L == 0:
                continue
            steps = 2 * L
            prev = bld[tr, tc] > 0.5
            cnt_in = 0
            cnt_w = 0
            for s in range(1, steps + 1):
                rr = int(round(tr + dr * s / steps))
                cc = int(round(tc + dc * s / steps))
                b = bld[rr, cc] > 0.5
                if b:
                    cnt_in += 1
                if b != prev:
                    cnt_w += 1
                prev = b
            d = np.sqrt(dr * dr + dc * dc)
            inside[r, c] = cnt_in * d / steps
            walls[r, c] = cnt_w
    return inside, walls


def physics_features(bld: np.ndarray, tx: tuple[int, int]):
    """Distance (m), metres of building on the straight Tx->Rx line, wall crossings, LOS flag."""
    rr, cc = np.mgrid[:bld.shape[0], :bld.shape[1]]
    d = np.hypot(rr - tx[0], cc - tx[1]).astype(np.float32)
    inside, walls = _line_features(bld, tx[0], tx[1])
    los = (inside == 0).astype(np.float32)
    return d, inside, walls, los
