"""Validation-only triangulated grid-post surfaces; both diagonals are explicit."""

from __future__ import annotations

import math

from mountain_twin.terrain.cell_reference import crossing_intervals
from mountain_twin.terrain.surface_reference import PostSurface, quadratic_maximum


class TriangulatedSurface(PostSurface):
    """Strict four-post support, shared with bilinear; no gap filling."""

    def __init__(self, dem, diagonal):
        if diagonal not in ("nw_se", "ne_sw"):
            raise ValueError("explicit nw_se or ne_sw diagonal required")
        super().__init__(dem)
        self.model = diagonal
        self.diagonal = diagonal

    def side(self, u, v):
        return u - v if self.diagonal == "nw_se" else 1 - u - v

    def plane(self, posts, upper):
        """Coefficients k,du,dv for z=k+du*u+dv*v; upper means side >= 0."""
        a, b, c, d = posts  # NW, NE, SW, SE
        if self.diagonal == "nw_se":
            return (a, b - a, d - b) if upper else (a, d - c, c - a)
        return (a, b - a, c - a) if upper else (b + c - d, d - c, d - b)

    def interpolate(self, posts, u, v):
        k, du, dv = self.plane(posts, self.side(u, v) >= 0)
        return k + du * u + dv * v


def triangle_intervals(surface, pixel_at, row, col, start, end):
    """Split a quad at its diagonal; short monotone diagonal coordinate required."""

    def side(d):
        u, v = pixel_at(d)
        return surface.side(u - col, v - row)

    checks = [side(start + (end - start) * i / 8) for i in range(9)]
    direction = 1 if checks[-1] >= checks[0] else -1
    if any(direction * (b - a) < -1e-10 for a, b in zip(checks, checks[1:])):
        raise ValueError("nonmonotone diagonal coordinate unsupported by local reference")
    if checks[0] * checks[-1] < 0:
        lo, hi = start, end
        for _ in range(48):
            mid = (lo + hi) / 2
            if direction * side(mid) < 0:
                lo = mid
            else:
                hi = mid
        root = (lo + hi) / 2
        # Match quad traversal event tolerance: sub-micrometre endpoint slivers
        # are indistinguishable from a diagonal starting/ending at the boundary.
        if min(root - start, end - root) <= 1e-7:
            return [(start, end)]
        return [(start, root), (root, end)]
    return [(start, end)]


def triangulated_horizon(surface, pixel_at, length_m, angle_tolerance_deg=1e-5):
    """Traverse quads AND triangles, then maximize each plane along the ray.

    Straight pixel-space rays give linear elevation. For geodesic curvature,
    quadratic local fitting and quarter-angle checks preserve the bilinear
    reference's numerical convention without treating the geodesic as straight.
    """
    if not isinstance(surface, TriangulatedSurface):
        raise ValueError("triangulated surface required")
    if not math.isfinite(angle_tolerance_deg) or angle_tolerance_deg <= 0:
        raise ValueError("positive finite angle tolerance 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),
        triangle_count=0,
        accepted_pieces=0,
        maximum_checked_residual_deg=0.0,
    )
    if origin.status != "valid":
        return result
    best = None
    missing = False
    for interval in intervals:
        quad, status = surface.quad(*pixel_at((interval.entry_m + interval.exit_m) / 2))
        if quad is None:
            missing = True
            continue
        row, col = quad
        posts, status = surface.posts(row, col)
        if posts is None:
            missing = True
            continue
        for start, end in triangle_intervals(
            surface, pixel_at, row, col, interval.entry_m, interval.exit_m
        ):
            result["triangle_count"] += 1
            u, v = pixel_at((start + end) / 2)
            upper = surface.side(u - col, v - row) >= 0
            k, du, dv = surface.plane(posts, upper)

            def elevation(d):
                u, v = pixel_at(d)
                return k + du * (u - col) + dv * (v - row)

            def solve(a, b, depth=0):
                width = b - a
                if width <= 0:
                    return math.degrees(math.atan2(elevation(b) - origin.elevation_m, b)), b
                z0 = elevation(a) - origin.elevation_m
                zm = elevation(a + width / 2) - origin.elevation_m
                z1 = elevation(b) - origin.elevation_m
                if a == 0:
                    z0 = 0.0
                q2 = 2 * (z1 - 2 * zm + z0) / width**2
                q1 = (z1 - z0) / width - q2 * width
                residual = 0.0
                for f in (0.25, 0.75):
                    t = f * width
                    actual = math.atan2(elevation(a + t) - origin.elevation_m, a + t)
                    fitted = math.atan2(q2 * t * t + q1 * t + z0, a + t)
                    residual = max(residual, abs(math.degrees(actual - fitted)))
                if residual > angle_tolerance_deg:
                    if depth >= 20:
                        raise ValueError(
                            f"triangle reference did not meet angle tolerance: {surface.diagonal} quad={row},{col} interval={start},{end} subinterval={a},{b} residual={residual}"
                        )
                    return max(
                        solve(a, a + width / 2, depth + 1),
                        solve(a + width / 2, b, 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(q2, q1, z0, a, width)

            angle, distance = solve(start, end)
            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=start,
                    exit_m=end,
                    positive_side=upper,
                )
    if not missing:
        result.update(horizon_angle_deg=best["angle_deg"], horizon_status="complete", maximum=best)
    return result
