import json
import os
from pathlib import Path

import numpy as np
import pytest
import torch

from rcm_ml.baseline import fit, predict_gray
from rcm_ml.data import DB_RANGE, N, PAPER_DB_PER_GRAY, P_TRNC, THRESH, gray_to_db, physics_features, split
from rcm_ml.train import Featurizer, Metrics

ROOT = Path(os.environ.get("RMS_ROOT", "data/rms"))
CACHE = Path(os.environ.get("RMS_CACHE", "data/rms_cache"))


def test_encode_decode_round_trip():
    assert gray_to_db(0.0) == pytest.approx(-147.0)
    assert gray_to_db(1.0) == pytest.approx(-47.84)
    db = np.linspace(-147, -47.84, 50)
    g = (db - P_TRNC) / DB_RANGE
    assert np.allclose(gray_to_db(g), db, atol=1e-4)
    # uint8 PNG round trip is exact to half a gray step (0.19 dB)
    u8 = np.round(g * 255).astype(np.uint8)
    assert np.abs(gray_to_db(u8 / 255.0) - db).max() <= DB_RANGE / 255 / 2 + 1e-4
    # the paper's "dB = 80 x gray" is the floor-clipped scale
    assert gray_to_db(THRESH) == pytest.approx(-127.2, abs=0.05)
    assert PAPER_DB_PER_GRAY == pytest.approx(79.3, abs=0.05)


def test_official_split_matches_radiounet_loader():
    ids = np.arange(0, 700, 1, dtype=np.int16)
    np.random.seed(42)
    np.random.shuffle(ids)
    tr, va, te = split(official=True)
    assert (len(tr), len(va), len(te)) == (501, 100, 99)
    assert (tr == ids[0:501] + 1).all() and (va == ids[501:601] + 1).all() and (te == ids[601:700] + 1).all()
    assert len(set(tr) | set(va) | set(te)) == 700 and 0 not in set(tr) | set(va) | set(te)


def _feat(coef, variant="feats", thresh=0.0):
    return Featurizer(torch.device("cpu"), coef, "B1", variant, N, thresh)


def test_tx_pixel_location_synthetic():
    b = {"bld": torch.zeros(2, N, N, dtype=torch.uint8), "y": torch.zeros(2, N, N, dtype=torch.uint8),
         "rc": torch.tensor([[10, 200], [128, 57]]), "inside": torch.zeros(2, N, N, dtype=torch.uint8),
         "walls": torch.zeros(2, N, N, dtype=torch.uint8)}
    x, *_ = _feat(np.zeros(5), "geom")(b)
    assert x.shape == (2, 2, N, N)
    for k, (r, c) in enumerate([(10, 200), (128, 57)]):
        assert x[k, 1].sum() == 1 and x[k, 1, r, c] == 1


@pytest.mark.skipif(not (ROOT / "antenna/1.json").exists(), reason="RadioMapSeer not downloaded")
def test_tx_pixel_matches_json():
    from rcm_ml.data import load_tx
    js = json.loads((ROOT / "antenna/1.json").read_text())
    for t in (0, 40, 79):
        r, c = load_tx(1, t)
        x, y = js[t]
        assert (r, c) == (255 - y, x)
        assert 53 <= r < 203 and 53 <= c < 203  # central 150x150


def test_device_b1_matches_numpy_b1():
    bld = np.zeros((N, N), np.float32)
    bld[100:140, 60:90] = 1
    bld[30:50, 150:230] = 1
    tx = (120, 128)
    _, inside, walls, _ = physics_features(bld, tx)
    coef = np.array([-66.0, -22.4, -0.17, -0.42, 13.3])
    ref = predict_gray(coef, inside, walls, *tx)
    b = {"bld": torch.from_numpy(bld.astype(np.uint8))[None], "y": torch.zeros(1, N, N, dtype=torch.uint8),
         "rc": torch.tensor([tx]), "inside": torch.from_numpy(np.clip(inside, 0, 255).astype(np.uint8))[None],
         "walls": torch.from_numpy(walls.astype(np.uint8))[None]}
    x, *_ = _feat(coef)(b)
    ins_q = np.clip(inside, 0, 255).astype(np.uint8).astype(np.float32)  # the cache stores whole metres
    ref_q = predict_gray(coef, ins_q, walls, *tx)
    assert np.abs(x[0, 2].numpy() - ref_q).max() < 1e-5
    # The cache truncates 'inside' to whole metres, so a line clipping a corner by < 1 m counts as LOS.
    # That changes B1 on only a few grazing pixels.
    assert (np.abs(ref_q - ref) > 1 / DB_RANGE).mean() < 0.02


def test_metrics_scales():
    y = torch.full((1, 1, 4, 4), 0.1)  # below the floor
    y[..., :2, :] = 0.5
    bld = torch.zeros_like(y)
    M = Metrics(THRESH)
    p_train = (y.clamp(min=THRESH) - THRESH) / (1 - THRESH)  # perfect prediction in the clipped scale
    M.add(p_train, y, bld)
    r = M.result()
    assert r["rmse_db_thr"] == pytest.approx(0, abs=1e-5) and r["rmse_db_above_-127"] == pytest.approx(0, abs=1e-5)
    assert r["rmse_db_raw"] == pytest.approx((0.1 * DB_RANGE) / 2 ** 0.5, rel=1e-4)  # half the pixels off by 0.1 gray


@pytest.mark.skipif(not (CACHE / "train_meta.json").exists(), reason="training cache not built")
def test_b1_fit_slope_is_physical():
    coef = fit(CACHE, "B1", "dpm")
    assert -25 < coef[1] < -19  # dB/decade; free space is -20, previous fit -22.4


def test_b2_geodesic_street_corner():
    from rcm_ml.b2 import b2_features, turns_from_cache
    b = np.ones((N, N), bool)
    b[100:110, 20:200] = False   # horizontal street
    b[100:250, 180:200] = False  # vertical street off its east end
    g, t = b2_features(b, (105, 30))
    assert g[105, 190] == 160 and turns_from_cache(t[105, 190]) == 0  # line of sight: Euclidean, no turn
    assert 270 < g[240, 190] < 290 and 0.8 < turns_from_cache(t[240, 190]) < 1.2  # one corner
    free = np.zeros((N, N), bool)
    g, t = b2_features(free, (128, 128))
    rr, cc = np.mgrid[:N, :N]
    assert (g == np.round(np.hypot(rr - 128, cc - 128))).all() and t.max() == 0


def test_device_b2_matches_numpy_b2():
    from rcm_ml.b2 import b2_features
    bld = np.zeros((N, N), np.float32)
    bld[100:140, 60:90] = 1
    bld[30:50, 150:230] = 1
    tx = (120, 128)
    _, inside, walls, _ = physics_features(bld, tx)
    geod, turns = b2_features(bld > 0, tx)
    iq, wq = np.clip(inside, 0, 255).astype(np.uint8), walls.astype(np.uint8)
    coef = np.array([-60.0, -25.0, -8.0, 10.0, -0.2, -5.0])
    ref = predict_gray(coef, iq.astype(np.float32), wq, *tx, name="B2", geod=geod, turns=turns)
    b = {"bld": torch.from_numpy(bld.astype(np.uint8))[None], "y": torch.zeros(1, N, N, dtype=torch.uint8),
         "rc": torch.tensor([tx]), "inside": torch.from_numpy(iq)[None], "walls": torch.from_numpy(wq)[None],
         "geod": torch.from_numpy(geod.astype(np.int32))[None], "turns": torch.from_numpy(turns)[None]}
    x, *_ = Featurizer(torch.device("cpu"), coef, "B2", "feats", N, 0.0)(b)
    assert np.abs(x[0, 2].numpy() - ref).max() < 1e-5
