"""Validation-only grid-post reconstruction and continuous quad horizon reference."""

from __future__ import annotations

import math

from mountain_twin.terrain.cell_reference import GEOD, cell_value, crossing_intervals
from mountain_twin.terrain.provider import Sample


class PostSurface:
    """Explicit nearest or bilinear surface over an existing GeoTiffDEM."""

    def __init__(self, dem, model="bilinear"):
        if model not in ("nearest", "bilinear"):
            raise ValueError("surface model must be nearest or bilinear")
        self.dem, self.model = dem, model
        self.inverse = ~dem.dataset.transform

    def pixels(self, latitude, longitude):
        if not (-90 <= latitude <= 90 and -180 <= longitude <= 180):
            raise ValueError("invalid WGS84 location")
        x, y = self.dem.transformer.transform(longitude, latitude, errcheck=True)
        t = self.inverse
        col, row = t.a * x + t.b * y + t.c, t.d * x + t.e * y + t.f
        if not all(math.isfinite(v) for v in (col, row)):
            raise ValueError("nonfinite transformed coordinates")
        return col - 0.5, row - 0.5

    def quad(self, col, row):
        h, w = self.dem.values.shape
        if not (0 <= col <= w - 1 and 0 <= row <= h - 1) or min(h, w) < 2:
            return None, "out_of_bounds"
        return (min(math.floor(row), h - 2), min(math.floor(col), w - 2)), "valid"

    def posts(self, row, col):
        values = [
            cell_value(self.dem, r, c)
            for r, c in ((row, col), (row, col + 1), (row + 1, col), (row + 1, col + 1))
        ]
        bad = next((status for _, status in values if status != "valid"), None)
        return (None, bad) if bad else ([z for z, _ in values], "valid")

    @staticmethod
    def interpolate(posts, u, v):
        a, b, c, d = posts
        return a * (1 - u) * (1 - v) + b * u * (1 - v) + c * (1 - u) * v + d * u * v

    def sample_post(self, col, row):
        if not all(math.isfinite(v) for v in (col, row)):
            raise ValueError("nonfinite post coordinates")
        if self.model == "nearest":
            z, status = cell_value(self.dem, math.floor(row + 0.5), math.floor(col + 0.5))
            return Sample(z, status)
        quad, status = self.quad(col, row)
        if quad is None:
            return Sample(None, status)
        r, c = quad
        posts, status = self.posts(r, c)
        return Sample(self.interpolate(posts, col - c, row - r) if posts else None, status)

    def sample(self, latitude, longitude):
        return self.sample_post(*self.pixels(latitude, longitude))


def quadratic_maximum(a, b, c, start, width):
    """Max (a*t²+b*t+c)/(start+t), t in [0,width], including right limit."""
    candidates = [0.0, width]
    # derivative numerator: a*t²+2*a*start*t+b*start-c
    if abs(a) > 1e-20:
        disc = start * start - (b * start - c) / a
        if disc >= 0:
            for t in (-start - math.sqrt(disc), -start + math.sqrt(disc)):
                if 0 < t < width:
                    candidates.append(t)

    def value(t):
        if start + t == 0:
            return b if abs(c) < 1e-10 else math.copysign(math.inf, c)
        return (a * t * t + b * t + c) / (start + t)

    t = max(candidates, key=value)
    return math.degrees(math.atan(value(t))), start + t


def continuous_horizon(surface, pixel_at, length_m, angle_tolerance_deg=1e-5, observer_height_m=0):
    """Bilinear reference on a monotone path parameterized by horizontal metres.

    pixel_at returns post (column,row) coordinates. Observer height lowers ray
    angles from the sampled terrain surface. Independent of production stepping;
    unknown support anywhere invalidates the whole horizon.
    """
    if surface.model != "bilinear":
        raise ValueError("continuous reference requires bilinear surface")
    if not math.isfinite(angle_tolerance_deg) or angle_tolerance_deg <= 0:
        raise ValueError("positive finite angle tolerance required")
    if not math.isfinite(observer_height_m) or observer_height_m < 0:
        raise ValueError("finite nonnegative observer height required")
    origin = surface.sample_post(*pixel_at(0))
    intervals = crossing_intervals(pixel_at, length_m)
    result = dict(
        observer_elevation_m=origin.elevation_m,
        horizon_angle_deg=None,
        horizon_status="incomplete",
        maximum=None,
        quad_count=len(intervals),
        accepted_pieces=0,
        maximum_checked_residual_deg=0.0,
    )
    if origin.status != "valid":
        return result
    best = None
    missing = False
    for interval in intervals:
        mid = (interval.entry_m + interval.exit_m) / 2
        quad, status = surface.quad(*pixel_at(mid))
        if quad is None:
            missing = True
            continue
        row, col = quad
        posts, status = surface.posts(row, col)
        if posts is None:
            missing = True
            continue

        def elevation(distance):
            u, v = pixel_at(distance)
            return surface.interpolate(posts, u - col, v - row)

        def solve(start, end, depth=0):
            width = end - start
            if width <= 0:
                return (
                    math.degrees(
                        math.atan2(
                            elevation(end) - origin.elevation_m - observer_height_m,
                            end,
                        )
                    ),
                    end,
                )
            z0 = elevation(start) - origin.elevation_m - observer_height_m
            zm = elevation(start + width / 2) - origin.elevation_m - observer_height_m
            z1 = elevation(end) - origin.elevation_m - observer_height_m
            a = 2 * (z1 - 2 * zm + z0) / (width * width)
            b = (z1 - z0) / width - a * width
            residual = 0.0
            for fraction in (0.25, 0.75):
                t = width * fraction
                actual = math.atan2(
                    elevation(start + t) - origin.elevation_m - observer_height_m,
                    start + t,
                )
                fitted = math.atan2(a * t * t + b * t + z0, start + t)
                residual = max(residual, abs(math.degrees(actual - fitted)))
            if residual > angle_tolerance_deg:
                if depth >= 20:
                    raise ValueError("continuous reference did not meet angle tolerance")
                return max(
                    solve(start, start + width / 2, depth + 1),
                    solve(start + width / 2, end, depth + 1),
                    key=lambda x: x[0],
                )
            result["accepted_pieces"] += 1
            result["maximum_checked_residual_deg"] = max(
                result["maximum_checked_residual_deg"], residual
            )
            return quadratic_maximum(a, b, z0, start, width)

        angle, distance = solve(interval.entry_m, interval.exit_m)
        if best is None or angle > best["angle_deg"]:
            best = dict(
                angle_deg=angle,
                distance_m=distance,
                quad_row=row,
                quad_col=col,
                posts_m=posts,
                entry_m=interval.entry_m,
                exit_m=interval.exit_m,
            )
    if not missing:
        result.update(horizon_angle_deg=best["angle_deg"], horizon_status="complete", maximum=best)
    return result


def geodesic_path(surface, latitude, longitude, azimuth):
    if not math.isfinite(azimuth):
        raise ValueError("finite azimuth required")

    def pixels(distance):
        lon, lat, _ = GEOD.fwd(longitude, latitude, azimuth, distance)
        return surface.pixels(lat, lon)

    return pixels
