"""Strict continuous bilinear sampling over aligned, native GeoTIFF grid-post tiles."""

from __future__ import annotations

import math
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable

import numpy as np
import rasterio
from pyproj import CRS, Transformer

from mountain_twin.terrain.cell_reference import crossing_intervals
from mountain_twin.terrain.provider import Sample, TerrainMetadata, vertical_reference


@dataclass(frozen=True)
class _Tile:
    path: Path
    values: np.ma.MaskedArray
    col0: int
    row0: int


class TiledGridPostSurface:
    """Aligned native grid posts with strict four-post bilinear support.

    GeoTIFF pixels are storage cells, while this helper treats their centers as
    explicitly declared elevation-post locations. It never fills a missing post
    or bridges an absent tile. ``quad``, ``posts`` and ``sample_post`` share the
    existing continuous-horizon reference interface.
    """

    def __init__(
        self,
        paths: Iterable[Path],
        *,
        source: str,
        product: str,
        vertical_datum: str,
        registration: str,
        provenance: tuple[tuple[str, str], ...] = (),
    ):
        if registration != "pixel_center_posts":
            raise ValueError("explicit pixel_center_posts registration required")
        paths = tuple(Path(path) for path in paths)
        if not paths:
            raise ValueError("at least one tile is required")
        self._datasets = []
        try:
            first = rasterio.open(paths[0])
            self._datasets.append(first)
            self.crs = CRS(first.crs) if first.crs else None
            if self.crs is None or not self.crs.is_projected:
                raise ValueError("tiles require one projected CRS")
            self.resolution = first.res
            if first.count != 1 or first.transform.b or first.transform.d:
                raise ValueError("tiles require one north-up band")
            if first.transform.a <= 0 or first.transform.e >= 0:
                raise ValueError("tiles require positive north-up resolution")
            if not math.isclose(first.res[0], first.res[1], abs_tol=1e-12):
                raise ValueError("tiles require square grid posts")
            if first.nodata is None or not math.isfinite(float(first.nodata)):
                raise ValueError("tiles require explicit finite nodata")
            self._nodata = float(first.nodata)
            self._scale, self._offset = first.scales[0], first.offsets[0]
            if not all(math.isfinite(v) for v in (self._scale, self._offset)):
                raise ValueError("nonfinite tile scale or offset")
            self._x0 = first.transform.c + first.transform.a / 2
            self._y0 = first.transform.f + first.transform.e / 2
            tiles = [self._read_tile(first, paths[0])]
            for path in paths[1:]:
                ds = rasterio.open(path)
                self._datasets.append(ds)
                if (
                    ds.count != 1
                    or CRS(ds.crs) != self.crs
                    or ds.res != self.resolution
                    or ds.transform.b != 0
                    or ds.transform.d != 0
                    or ds.transform.a <= 0
                    or ds.transform.e >= 0
                    or ds.nodata is None
                    or not math.isclose(float(ds.nodata), self._nodata, abs_tol=0)
                    or ds.scales[0] != self._scale
                    or ds.offsets[0] != self._offset
                ):
                    raise ValueError("inconsistent tile metadata")
                tiles.append(self._read_tile(ds, path))
            self._tiles = tuple(tiles)
            self.transformer = Transformer.from_crs("EPSG:4326", self.crs, always_xy=True)
            bounds = [ds.bounds for ds in self._datasets]
            self.info = TerrainMetadata(
                source=source,
                product=product,
                crs=self.crs.to_string(),
                bounds=(
                    min(b.left for b in bounds),
                    min(b.bottom for b in bounds),
                    max(b.right for b in bounds),
                    max(b.top for b in bounds),
                ),
                horizontal_resolution=(self.resolution[0], self.resolution[1]),
                horizontal_units=("metre", "metre"),
                elevation_unit="m",
                vertical_reference=vertical_reference(vertical_datum),
                vertical_reference_verified=vertical_reference(vertical_datum) is not None,
                sampling="strict_bilinear_pixel_center_posts",
                nodata_policy="any required missing post -> nodata; absent tile -> out_of_bounds",
                provenance=provenance,
            )
            self.model = "bilinear"
        except Exception:
            self.close()
            raise

    def _read_tile(self, ds, path):
        col0_float = (ds.transform.c + ds.transform.a / 2 - self._x0) / self.resolution[0]
        row0_float = (self._y0 - (ds.transform.f + ds.transform.e / 2)) / self.resolution[1]
        col0, row0 = round(col0_float), round(row0_float)
        if not (
            math.isclose(col0_float, col0, abs_tol=1e-8)
            and math.isclose(row0_float, row0, abs_tol=1e-8)
        ):
            raise ValueError("tile grid posts are not aligned")
        return _Tile(path, ds.read(1, masked=True), col0, row0)

    def close(self):
        for ds in getattr(self, "_datasets", []):
            ds.close()
        self._datasets = []

    def __enter__(self):
        return self

    def __exit__(self, *args):
        self.close()

    def pixels(self, latitude: float, longitude: float):
        if not (-90 <= latitude <= 90 and -180 <= longitude <= 180):
            raise ValueError("invalid WGS84 location")
        x, y = self.transformer.transform(longitude, latitude, errcheck=True)
        return self.projected_pixels(x, y)

    def projected_pixels(self, x: float, y: float):
        if not all(math.isfinite(v) for v in (x, y)):
            raise ValueError("nonfinite projected coordinates")
        return (x - self._x0) / self.resolution[0], (self._y0 - y) / self.resolution[1]

    def _post(self, col: int, row: int):
        for tile in self._tiles:
            local_col, local_row = col - tile.col0, row - tile.row0
            if 0 <= local_row < tile.values.shape[0] and 0 <= local_col < tile.values.shape[1]:
                value = tile.values[local_row, local_col]
                if np.ma.is_masked(value) or not math.isfinite(float(value)):
                    return None, "nodata"
                z = float(value) * self._scale + self._offset
                return (z, "valid") if math.isfinite(z) else (None, "nodata")
        return None, "out_of_bounds"

    def quad(self, col: float, row: float):
        if not all(math.isfinite(v) for v in (col, row)):
            raise ValueError("nonfinite grid coordinates")
        base_col, base_row = math.floor(col), math.floor(row)
        candidates = [(base_col, base_row)]
        if math.isclose(col, base_col, abs_tol=1e-10):
            candidates.append((base_col - 1, base_row))
        if math.isclose(row, base_row, abs_tol=1e-10):
            candidates.append((base_col, base_row - 1))
        if math.isclose(col, base_col, abs_tol=1e-10) and math.isclose(
            row, base_row, abs_tol=1e-10
        ):
            candidates.append((base_col - 1, base_row - 1))
        statuses = []
        for candidate in candidates:
            _, status = self.posts(candidate[1], candidate[0])
            if status == "valid":
                return (candidate[1], candidate[0]), "valid"
            # A valid neighbouring quad must not be used to bypass a nodata post
            # in the canonical interpolation support at a grid boundary.
            if status == "nodata":
                return None, status
            statuses.append(status)
        return None, "nodata" if "nodata" in statuses else "out_of_bounds"

    def posts(self, row: int, col: int):
        values = [
            self._post(c, r)
            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: float, v: float):
        nw, ne, sw, se = posts
        return nw * (1 - u) * (1 - v) + ne * u * (1 - v) + sw * (1 - u) * v + se * u * v

    def sample_post(self, col: float, row: float):
        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), "valid")

    def sample_projected(self, x: float, y: float):
        return self.sample_post(*self.projected_pixels(x, y))

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


def native_bilinear_horizon(
    surface: TiledGridPostSurface,
    pixel_at,
    length_m: float,
    position_tolerance_m: float = 1e-6,
    observer_height_m: float = 0,
    observer_elevation_m: float | None = None,
):
    """Maximize a tiled bilinear surface on every crossed native grid quad.

    Grid-boundary traversal defines the pieces; a bounded golden-section search
    maximizes the angle on each piece. It does not use an arbitrary ray step.
    ``observer_elevation_m`` explicitly supplies an absolute observer elevation
    for hybrid geometries whose observer and occluding surfaces differ. It is
    mutually exclusive with ``observer_height_m``. At zero same-surface height,
    the one-sided directional-slope limit is a candidate; otherwise its limit is
    -90 degrees.
    """
    if surface.model != "bilinear":
        raise ValueError("bilinear tiled surface required")
    if not math.isfinite(length_m) or length_m <= 0:
        raise ValueError("positive finite ray length required")
    if not math.isfinite(position_tolerance_m) or position_tolerance_m <= 0:
        raise ValueError("positive finite position tolerance required")
    if not math.isfinite(observer_height_m) or observer_height_m < 0:
        raise ValueError("finite nonnegative observer height required")
    if observer_elevation_m is not None and (
        not math.isfinite(observer_elevation_m) or observer_height_m != 0
    ):
        raise ValueError("explicit observer elevation requires zero finite observer height")
    origin = surface.sample_post(*pixel_at(0))
    observer_z = (
        observer_elevation_m
        if observer_elevation_m is not None
        else (origin.elevation_m + observer_height_m if origin.elevation_m is not None else None)
    )
    result = dict(
        observer_elevation_m=observer_z,
        horizon_angle_deg=None,
        horizon_status="incomplete",
        maximum=None,
        quad_count=0,
        evaluated_intervals=0,
        position_tolerance_m=position_tolerance_m,
    )
    if origin.status != "valid":
        return result
    if observer_elevation_m is not None:
        result["observer_surface_elevation_m"] = origin.elevation_m
    best = None
    missing = False
    for interval in crossing_intervals(pixel_at, length_m):
        if interval.exit_m <= interval.entry_m:
            continue
        result["quad_count"] += 1
        midpoint = (interval.entry_m + interval.exit_m) / 2
        quad, status = surface.quad(*pixel_at(midpoint))
        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 angle(distance, slope_limit=None):
            if distance == 0:
                return math.degrees(math.atan(slope_limit))
            return math.degrees(
                math.atan2(
                    elevation(distance) - observer_z,
                    distance,
                )
            )

        start, end = interval.entry_m, interval.exit_m
        candidates = [(angle(end), end)]
        if start > 0:
            candidates.append((angle(start), start))
        elif observer_elevation_m is None and observer_height_m == 0:
            epsilon = min(end / 1024, 1e-5)
            slope = (elevation(epsilon) - observer_z) / epsilon
            candidates.append((angle(0, slope), 0.0))
            start = epsilon

        # Four independent bounded searches protect against a local non-unimodal
        # composition of the projected geodesic and bilinear grid surface.
        for left, right in zip(
            [start + (end - start) * part / 4 for part in range(4)],
            [start + (end - start) * part / 4 for part in range(1, 5)],
        ):
            a, b = left, right
            ratio = (math.sqrt(5) - 1) / 2
            c, d = b - ratio * (b - a), a + ratio * (b - a)
            fc, fd = angle(c), angle(d)
            while b - a > position_tolerance_m:
                if fc >= fd:
                    b, d, fd = d, c, fc
                    c = b - ratio * (b - a)
                    fc = angle(c)
                else:
                    a, c, fc = c, d, fd
                    d = a + ratio * (b - a)
                    fd = angle(d)
            candidates.append(max(((fc, c), (fd, d)), key=lambda item: item[0]))
            result["evaluated_intervals"] += 1
        candidate_angle, candidate_distance = max(candidates, key=lambda item: item[0])
        if best is None or candidate_angle > best["angle_deg"]:
            best = dict(
                angle_deg=candidate_angle,
                distance_m=candidate_distance,
                quad_row=row,
                quad_col=col,
                posts_m=posts,
                entry_m=interval.entry_m,
                exit_m=interval.exit_m,
                elevation_m=elevation(candidate_distance)
                if candidate_distance > 0
                else origin.elevation_m,
            )
    if not missing:
        result.update(horizon_angle_deg=best["angle_deg"], horizon_status="complete", maximum=best)
    return result
