import os
from pathlib import Path

import numpy as np
import pytest

from rcm_ml.raypaths import EXIT, REFLECT, STOP, deposit_paths, pad_paths, ray_path, scene, trace_rays, tx_xy
from rcm_ml.raytrace import EPS_C, LAM, _shoot

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


def _box():
    # one 20 x 20 m building in the middle of the tile, in (col, row) coordinates
    return [np.array([(100.0, 100.0), (120.0, 100.0), (120.0, 120.0), (100.0, 120.0)])]


def test_ray_path_reflects_specularly_and_exits():
    free, segs, _ = scene(_box())
    # ray along +x at y = 110 from x = 50 hits the left wall (x = 100) head-on, comes back and leaves at x = 0
    p, k, w, c, l = ray_path(segs, 50.0, 110.0, 1.0, 0.0, 2, 256, EPS_C.real, EPS_C.imag, 10.0)
    assert list(k) == [REFLECT, EXIT]
    assert np.allclose(p, [[50, 110], [100, 110], [0, 110]])
    assert c[0] == pytest.approx(1.0)
    assert l[1] == pytest.approx(0.1)
    # with 0 bounces a ray stops at the first wall
    p, k, *_ = ray_path(segs, 90.0, 95.0, np.sqrt(0.5), np.sqrt(0.5), 0, 256, EPS_C.real, EPS_C.imag, 10.0)
    assert list(k) == [STOP] and np.allclose(p[-1], [100, 105], atol=1e-6)


def _reference(free, segs, tx, n_rays, bounces, step=0.5, ple=2.6):
    acc = np.zeros(free.shape)
    dth = 2 * np.pi / n_rays
    for k in range(n_rays):
        a = (k + 0.5) * dth
        _shoot(acc, free, segs, tx[0], tx[1], np.cos(a), np.sin(a), 0.0, 1.0, dth, bounces, LAM, EPS_C.real, EPS_C.imag, -1, 0.0, 0.0, step, ple, -1.0, 10.0)
    return acc


def _check(polys, tx_rc):
    free, segs, _ = scene(polys)
    tx = tx_xy(tx_rc)
    n_rays, B = 720, 2
    ang = (np.arange(n_rays) + 0.5) * 2 * np.pi / n_rays
    pts, kind, wall, cosi, lin = pad_paths(trace_rays(segs, tx, ang, B), B)
    ours = deposit_paths(free, pts, kind, lin, n_rays, ple=2.6)
    ref = _reference(free, segs, tx, n_rays, B)
    m = (ref > 0) | (ours > 0)
    assert m.sum() > 1000
    db = 10 * np.log10(np.maximum(ours[m], 1e-30)) - 10 * np.log10(np.maximum(ref[m], 1e-30))
    # identical up to the last half step at the tile edge
    assert np.mean(np.abs(db) < 1e-3) > 0.999


def test_deposit_matches_tracer_box():
    _check(_box(), (60, 60))


@pytest.mark.skipif(not (ROOT / "polygon/buildings_complete/1.json").exists(), reason="RadioMapSeer not downloaded")
def test_deposit_matches_tracer_real_map():
    from rcm_ml.data import load_tx
    from rcm_ml.raytrace import load_polygons
    _check(load_polygons(1, ROOT), load_tx(1, 0))
