"""Analytic local-metric surfaces for tests, never an observed DEM product."""

from __future__ import annotations

import math
from typing import Callable

from pyproj import CRS, Transformer

from mountain_twin.terrain.provider import Sample, TerrainMetadata


class SyntheticTerrain:
    """Evaluate a pure function of local east/north metres in an azimuthal grid.

    A surface returning None or nonfinite elevation represents nodata. Bounds are
    west/north inclusive, east/south exclusive, as for the raster adapter.
    Callers must supply a deterministic function and a descriptive surface ID.
    """

    def __init__(
        self,
        surface: Callable[[float, float], float | None],
        *,
        surface_id: str,
        latitude: float = 0,
        longitude: float = 0,
        bounds: tuple[float, float, float, float] = (-10000, -10000, 10000, 10000),
    ):
        if not surface_id.strip():
            raise ValueError("surface_id is required for provenance")
        if not (-90 < latitude < 90 and -180 <= longitude <= 180):
            raise ValueError("invalid synthetic origin")
        if not all(math.isfinite(v) for v in bounds) or not (
            bounds[0] < bounds[2] and bounds[1] < bounds[3]
        ):
            raise ValueError("invalid synthetic bounds")
        crs = CRS.from_proj4(
            f"+proj=aeqd +lat_0={latitude} +lon_0={longitude} +datum=WGS84 +units=m +no_defs"
        )
        self._transform = Transformer.from_crs(4326, crs, always_xy=True, allow_ballpark=False)
        self._surface = surface
        self.info = TerrainMetadata(
            source="synthetic analytic surface",
            product=surface_id,
            crs=crs.to_string(),
            bounds=bounds,
            horizontal_resolution=None,
            horizontal_units=("metre", "metre"),
            elevation_unit="m",
            vertical_reference=None,
            vertical_reference_verified=False,
            sampling="analytic function in local AEQD metres",
            nodata_policy="None/nonfinite -> nodata; no filling",
            provenance=(("surface_id", surface_id), ("origin", f"{latitude},{longitude}")),
        )

    def sample(self, latitude: float, longitude: float) -> Sample:
        if not (-90 <= latitude <= 90 and -180 <= longitude <= 180):
            raise ValueError("invalid WGS84 sample coordinates")
        x, y = self._transform.transform(longitude, latitude, errcheck=True)
        if not all(math.isfinite(v) for v in (x, y)):
            raise ValueError("nonfinite transformed coordinates")
        west, south, east, north = self.info.bounds
        if not (west <= x < east and south < y <= north):
            return Sample(None, "out_of_bounds")
        value = self._surface(x, y)
        if value is None or not math.isfinite(value):
            return Sample(None, "nodata")
        return Sample(float(value), "valid")
