"""GeoTIFF adapter. Horizontal reprojection only; no vertical datum conversion."""

from __future__ import annotations

import hashlib
import math
from pathlib import Path

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

# Preserve existing import paths for callers of the original prototype.
from mountain_twin.terrain.provider import ElevationSource as ElevationSource
from mountain_twin.terrain.provider import Sample, TerrainMetadata, vertical_reference


class GeoTiffDEM:
    """Single-band north-up raster, nearest-cell sampling, explicit source contract."""

    def __init__(self, path: Path, metadata: dict):
        for key in ("source", "vertical_datum", "crs", "elevation_unit"):
            if not isinstance(metadata.get(key), str) or not metadata[key].strip():
                raise ValueError(f"DEM metadata requires nonempty {key}")
        if metadata["elevation_unit"] != "m":
            raise ValueError("only explicitly declared metre elevations are supported")
        self.dataset = rasterio.open(path)
        try:
            ds = self.dataset
            if ds.driver != "GTiff" or ds.count != 1:
                raise ValueError("DEM must be a single-band GeoTIFF")
            if ds.crs is None:
                raise ValueError("DEM CRS is missing")
            crs = CRS(ds.crs)
            if not (crs.is_projected or crs.is_geographic) or crs.is_compound:
                raise ValueError("DEM requires a horizontal geographic or projected CRS")
            if crs != CRS(metadata["crs"]):
                raise ValueError("DEM CRS mismatch with declared metadata")
            t = ds.transform
            if not all(math.isfinite(x) for x in (*t, *ds.bounds)):
                raise ValueError("nonfinite transform or bounds")
            if t.b != 0 or t.d != 0 or t.a <= 0 or t.e >= 0:
                raise ValueError("prototype requires a north-up raster with positive resolution")
            if ds.width <= 0 or ds.height <= 0:
                raise ValueError("empty DEM bounds")
            if ds.units[0] not in (None, "m", "metre", "meter"):
                raise ValueError("raster elevation unit conflicts with metre declaration")
            self.transformer = Transformer.from_crs(
                "EPSG:4326", crs, always_xy=True, allow_ballpark=False
            )
            self.values = ds.read(1, masked=True)
            self.scale, self.offset = ds.scales[0], ds.offsets[0]
            if not all(math.isfinite(x) for x in (self.scale, self.offset)):
                raise ValueError("nonfinite raster scale/offset")
            digest = hashlib.sha256()
            with path.open("rb") as handle:
                for block in iter(lambda: handle.read(1024 * 1024), b""):
                    digest.update(block)
            self.metadata = dict(
                metadata,
                sha256=digest.hexdigest(),
                raster_crs=crs.to_string(),
                bounds=list(ds.bounds),
                resolution=list(ds.res),
                resolution_units=[axis.unit_name for axis in crs.axis_info[:2]],
                width=ds.width,
                height=ds.height,
                nodata=str(ds.nodata),
                scale=self.scale,
                offset=self.offset,
                sampling="nearest_cell",
                scientific_validation="source and vertical datum require independent validation",
            )
            reference = vertical_reference(metadata["vertical_datum"])
            self.info = TerrainMetadata(
                source=metadata["source"],
                product=metadata.get("product"),
                crs=crs.to_string(),
                bounds=tuple(ds.bounds),
                horizontal_resolution=tuple(ds.res),
                horizontal_units=tuple(axis.unit_name for axis in crs.axis_info[:2]),
                elevation_unit="m",
                vertical_reference=reference,
                vertical_reference_verified=(
                    reference is not None and metadata.get("vertical_reference_verified") is True
                ),
                sampling="nearest_cell",
                nodata_policy="masked/nonfinite -> nodata; no filling",
                provenance=(
                    ("sha256", digest.hexdigest()),
                    (
                        "vertical_reference_evidence",
                        str(metadata.get("vertical_reference_evidence", "not supplied")),
                    ),
                ),
            )
        except Exception:
            self.dataset.close()
            raise

    def __enter__(self):
        return self

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

    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.transformer.transform(longitude, latitude, errcheck=True)
        if not all(math.isfinite(v) for v in (x, y)):
            raise ValueError("nonfinite transformed coordinates")
        row, col = self.dataset.index(x, y)
        if not (0 <= row < self.dataset.height and 0 <= col < self.dataset.width):
            return Sample(None, "out_of_bounds")
        value = self.values[row, col]
        if np.ma.is_masked(value) or not math.isfinite(float(value)):
            return Sample(None, "nodata")
        elevation = float(value) * self.scale + self.offset
        if not math.isfinite(elevation):
            return Sample(None, "nodata")
        return Sample(elevation, "valid")
