"""Individual ray trajectories from our 2D tracer (rcm_ml/raytrace.py), the teacher for the learned ray tracer.

`trace` / `radio_map` only accumulate power per pixel. Here one ray from the Tx is followed exactly like
`raytrace._shoot` (same `_hit`, same self-wall skip, same specular reflection) and its polyline is recorded:

  P0 = Tx, P1, ..., Pk     points; leg i goes from P(i-1) to P(i)
  kind[i]                  how leg i ends: REFLECT (hit a wall and reflects), EXIT (leaves the 256 m tile),
                           STOP (hit a wall after the last allowed bounce)
  wall[i]                  index of the wall segment hit (-1 for EXIT)
  cos_i[i]                 incidence cosine at that wall
  lin[i]                   power factor carried *into* leg i (product of reflection losses before it)

A ray leaving the tile stops: `_deposit` stops at the tile edge, and a leg starting outside deposits nothing, so
reflections outside the tile never reach a pixel. `deposit_paths` turns recorded polylines back into the tracer's
power map with `_deposit`; tests/test_raypaths.py checks it reproduces the LOS/R/RR part of `trace`.
"""
from __future__ import annotations

import numba as nb
import numpy as np

from .raytrace import EPS_C, LAM, _deposit, _gamma2, _hit, geometry, rasterize

REFLECT, EXIT, STOP = 0, 1, 2


@nb.njit(cache=True)
def _exit_t(ox, oy, dx, dy, n):
    """Distance along o + t d to the boundary of the [0, n]^2 tile (o inside)."""
    t = np.inf
    if dx > 1e-12:
        t = min(t, (n - ox) / dx)
    elif dx < -1e-12:
        t = min(t, -ox / dx)
    if dy > 1e-12:
        t = min(t, (n - oy) / dy)
    elif dy < -1e-12:
        t = min(t, -oy / dy)
    return t


@nb.njit(cache=True)
def ray_path(segs, ox, oy, dx, dy, bounces, n, eps_re, eps_im, rloss):
    """One ray, up to `bounces` reflections. Returns (pts (k+1, 2), kind (k,), wall (k,), cos_i (k,), lin (k,))."""
    pts = np.zeros((bounces + 2, 2))
    kind = np.zeros(bounces + 1, np.int64)
    wall = np.full(bounces + 1, -1, np.int64)
    cosi = np.zeros(bounces + 1)
    lin = np.ones(bounces + 1)
    pts[0, 0], pts[0, 1] = ox, oy
    skip = -1
    lin0 = 1.0
    k = 0
    for b in range(bounces + 1):
        t, i = _hit(ox, oy, dx, dy, segs, skip, -1)
        te = _exit_t(ox, oy, dx, dy, n)
        lin[b] = lin0
        if i < 0 or te < t:
            pts[b + 1, 0], pts[b + 1, 1] = ox + te * dx, oy + te * dy
            kind[b] = EXIT
            k = b + 1
            break
        ex, ey = segs[i, 2] - segs[i, 0], segs[i, 3] - segs[i, 1]
        L = np.hypot(ex, ey)
        nx, ny = -ey / L, ex / L
        dot = dx * nx + dy * ny
        ox, oy = ox + t * dx, oy + t * dy
        pts[b + 1, 0], pts[b + 1, 1] = ox, oy
        wall[b] = i
        cosi[b] = abs(dot)
        k = b + 1
        if b == bounces:
            kind[b] = STOP
            break
        kind[b] = REFLECT
        lin0 *= _gamma2(abs(dot), eps_re, eps_im) if rloss < 0 else 10.0 ** (-rloss / 10.0)
        dx, dy = dx - 2 * dot * nx, dy - 2 * dot * ny
        skip = i
    return pts[: k + 1], kind[:k], wall[:k], cosi[:k], lin[:k]


def scene(polys, n=256):
    """(free mask, wall segments, convex corners) of a tile, as the tracer builds them."""
    free = ~rasterize(polys, n)
    segs, corners = geometry(polys, free)
    return free, segs, corners


def tx_xy(tx_rc):
    """Tracer coordinates (x = col + 0.5, y = row + 0.5) of a Tx given as (row, col) pixel."""
    return tx_rc[1] + 0.5, tx_rc[0] + 0.5


def trace_rays(segs, tx, angles, bounces=2, n=256, eps=EPS_C, rloss=10.0):
    """Polylines of rays launched from tx = (x, y) at `angles` (radians, x towards +col, y towards +row).
    rloss: dB per reflection (calibrated tracer: 10), < 0 for Fresnel concrete."""
    out = []
    for a in angles:
        out.append(ray_path(segs, tx[0], tx[1], np.cos(a), np.sin(a), bounces, n, eps.real, eps.imag, rloss))
    return out


def pad_paths(paths, bounces):
    """Fixed-size arrays: pts (R, B+2, 2), kind (R, B+1) with -1 = no leg, wall, cos_i, lin."""
    R, K = len(paths), bounces + 1
    pts = np.full((R, K + 1, 2), np.nan, np.float32)
    kind = np.full((R, K), -1, np.int8)
    wall = np.full((R, K), -1, np.int32)
    cosi = np.zeros((R, K), np.float32)
    lin = np.zeros((R, K), np.float32)
    for r, (p, k, w, c, l) in enumerate(paths):
        pts[r, : len(p)] = p
        kind[r, : len(k)] = k
        wall[r, : len(w)] = w
        cosi[r, : len(c)] = c
        lin[r, : len(l)] = l
    return pts, kind, wall, cosi, lin


@nb.njit(cache=True)
def _deposit_polylines(acc, free, pts, kind, lin, dth, lam, step, ple):
    for r in range(pts.shape[0]):
        d0 = 0.0
        for i in range(kind.shape[1]):
            if kind[r, i] < 0:
                break
            ox, oy = pts[r, i, 0], pts[r, i, 1]
            vx, vy = pts[r, i + 1, 0] - ox, pts[r, i + 1, 1] - oy
            L = np.hypot(vx, vy)
            if L < 1e-9:
                break
            dx, dy = vx / L, vy / L
            # the tracer clips every leg at 2 n; legs end inside the tile here, so L < 2 n always
            _deposit(acc, free, ox, oy, dx, dy, min(L, 2.0 * acc.shape[0]), d0, d0, lin[r, i], dth, lam, 0.0, 0.0, step, ple, -1.0, -1.0)
            d0 += L


def deposit_paths(free, pts, kind, lin, n_rays, step=0.5, ple=2.6049672462988367):
    """Linear path gain per pixel from padded polylines (the tracer's track-length estimator, ray spacing 2 pi / n_rays).
    `lin` is the power carried into each leg. Same as trace(...)'s LOS/R/RR part when the paths are the teacher's."""
    acc = np.zeros(free.shape)
    _deposit_polylines(acc, free, pts.astype(np.float64), kind.astype(np.int64), lin.astype(np.float64), 2 * np.pi / n_rays, LAM, step, ple)
    return acc
