"""Validation-only exact boundary traversal of a monotone pixel-coordinate ray.

Reference surface: piecewise-constant nearest-cell footprints, not interpolated
DSM posts or physical terrain. See docs/near_field_geometry.md for conventions.
"""

from __future__ import annotations

import math
from dataclasses import asdict, dataclass
from typing import Callable

import numpy as np
from pyproj import Geod

GEOD = Geod(ellps="WGS84")


@dataclass(frozen=True)
class CellInterval:
    row: int
    col: int
    entry_m: float
    exit_m: float


def crossing_intervals(pixel_at: Callable[[float], tuple[float, float]], length_m: float):
    """Enumerate crossed cells by integer grid-boundary root intersections.

    Coordinates are (column,row), in a monotone local ray. No elevations or fixed
    terrain step are used. Endpoint-only cells follow the same floor convention.
    """
    if not math.isfinite(length_m) or length_m <= 0:
        raise ValueError("positive finite ray length required")
    checkpoints = [pixel_at(length_m * i / 128) for i in range(129)]
    if not all(math.isfinite(v) for p in checkpoints for v in p):
        raise ValueError("nonfinite pixel coordinates")
    start, end = checkpoints[0], checkpoints[-1]
    events = [0.0, length_m]
    for axis in (0, 1):
        direction = 1 if end[axis] >= start[axis] else -1
        if any(
            direction * (b[axis] - a[axis]) < -1e-10 for a, b in zip(checkpoints, checkpoints[1:])
        ):
            raise ValueError("reference requires monotone pixel ray")
        lower, upper = sorted((start[axis], end[axis]))
        if upper - lower > 100000:
            raise ValueError("too many cell boundaries")
        for line in range(math.floor(lower) + 1, math.ceil(upper)):
            lo, hi = 0.0, length_m
            for _ in range(48):
                middle = (lo + hi) / 2
                if direction * (pixel_at(middle)[axis] - line) < 0:
                    lo = middle
                else:
                    hi = middle
            events.append((lo + hi) / 2)
    unique = []
    for event in sorted(events):
        if not unique or event - unique[-1] > 1e-7:
            unique.append(event)
    unique[-1] = length_m
    result = []
    for entry, exit_distance in zip(unique, unique[1:]):
        col, row = pixel_at((entry + exit_distance) / 2)
        result.append(CellInterval(math.floor(row), math.floor(col), entry, exit_distance))
    col, row = pixel_at(length_m)
    if (math.floor(row), math.floor(col)) != (result[-1].row, result[-1].col):
        result.append(CellInterval(math.floor(row), math.floor(col), length_m, length_m))
    return result


def footprint_angle(delta_m, entry_m, exit_m):
    """Supremum angle over a constant-height ray interval."""
    distance = entry_m if delta_m > 0 else exit_m
    return math.degrees(math.atan2(delta_m, distance)), distance


def cell_value(dem, row, col):
    if not (0 <= row < dem.dataset.height and 0 <= col < dem.dataset.width):
        return None, "out_of_bounds"
    value = dem.values[row, col]
    if np.ma.is_masked(value) or not math.isfinite(float(value)):
        return None, "nodata"
    value = float(value) * dem.scale + dem.offset
    return (value, "valid") if math.isfinite(value) else (None, "nodata")


def reference_horizon(dem, latitude, longitude, azimuth_deg, length_m=5000, observer_height_m=0):
    if (
        not all(math.isfinite(v) for v in (latitude, longitude, azimuth_deg, observer_height_m))
        or observer_height_m < 0
    ):
        raise ValueError("finite inputs and nonnegative height required")
    if not (-90 <= latitude <= 90 and -180 <= longitude <= 180):
        raise ValueError("invalid WGS84 location")
    inverse = ~dem.dataset.transform

    def pixels(distance):
        lon, lat, _ = GEOD.fwd(longitude, latitude, azimuth_deg, distance)
        x, y = dem.transformer.transform(lon, lat, errcheck=True)
        return (
            inverse.a * x + inverse.b * y + inverse.c,
            inverse.d * x + inverse.e * y + inverse.f,
        )

    col, row = pixels(0)
    origin = (math.floor(row), math.floor(col))
    origin_z, origin_status = cell_value(dem, *origin)
    cells = []
    for interval in crossing_intervals(pixels, length_m):
        z, status = cell_value(dem, interval.row, interval.col)
        angle = distance = None
        if status == "valid" and origin_status == "valid":
            angle, distance = footprint_angle(
                z - origin_z - observer_height_m, interval.entry_m, interval.exit_m
            )
        cells.append(
            dict(
                **asdict(interval),
                elevation_m=z,
                status=status,
                angle_deg=angle,
                angle_distance_m=distance,
                is_origin=(interval.row, interval.col) == origin,
            )
        )
    complete = origin_status == "valid" and all(c["status"] == "valid" for c in cells)
    maximum = max(cells, key=lambda c: c["angle_deg"]) if complete else None
    nonorigin = [c for c in cells if not c["is_origin"]]
    excluded = max(nonorigin, key=lambda c: c["angle_deg"]) if complete and nonorigin else None
    return dict(
        origin_row=origin[0],
        origin_col=origin[1],
        origin_subcell_col=col - math.floor(col),
        origin_subcell_row=row - math.floor(row),
        origin_elevation_m=origin_z,
        horizon_angle_deg=maximum["angle_deg"] if maximum else None,
        horizon_status="complete" if complete else "incomplete",
        maximum_cell=maximum,
        origin_excluded_horizon_deg=excluded["angle_deg"] if excluded else None,
        cells=cells,
    )
