"""Lossless native-grid crop/assembly. Missing coverage remains explicit nodata."""

import math

import numpy as np
from rasterio.windows import Window, from_bounds, transform


def native_crop(datasets, bounds):
    """Return masked native-grid pixels and transform, rejecting grid/overlap conflicts.

    No interpolation, reprojection or averaging. Extent snaps outward on the
    first input's grid. Valid overlapping values must agree exactly.
    """
    if not datasets or not all(math.isfinite(v) for v in bounds):
        raise ValueError("datasets and finite bounds required")
    if not (bounds[0] < bounds[2] and bounds[1] < bounds[3]):
        raise ValueError("invalid bounds")
    first = datasets[0]
    t = first.transform
    if t.b or t.d or t.a <= 0 or t.e >= 0:
        raise ValueError("north-up grid required")
    window = from_bounds(*bounds, transform=t)
    left, top = math.floor(window.col_off), math.floor(window.row_off)
    right = math.ceil(window.col_off + window.width)
    bottom = math.ceil(window.row_off + window.height)
    output = np.ma.masked_all((bottom - top, right - left), dtype=first.dtypes[0])
    for ds in datasets:
        if (
            ds.crs != first.crs
            or ds.count != 1
            or ds.dtypes != first.dtypes
            or ds.scales != first.scales
            or ds.offsets != first.offsets
            or ds.units != first.units
            or ds.transform.b
            or ds.transform.d
            or not np.allclose(ds.res, first.res, rtol=0, atol=1e-12)
        ):
            raise ValueError("incompatible native grids")
        col_offset = (ds.transform.c - t.c) / t.a
        row_offset = (ds.transform.f - t.f) / t.e
        if not np.allclose(
            [col_offset, row_offset], np.round([col_offset, row_offset]), rtol=0, atol=1e-7
        ):
            raise ValueError("unaligned native grids")
        col_offset, row_offset = round(col_offset), round(row_offset)
        intersection_left, intersection_right = (
            max(left, col_offset),
            min(right, col_offset + ds.width),
        )
        u, b = max(top, row_offset), min(bottom, row_offset + ds.height)
        if intersection_left >= intersection_right or u >= b:
            continue
        src = ds.read(
            1,
            window=Window(
                intersection_left - col_offset,
                u - row_offset,
                intersection_right - intersection_left,
                b - u,
            ),
            masked=True,
        )
        src = np.ma.masked_invalid(src)
        dst = output[u - top : b - top, intersection_left - left : intersection_right - left]
        valid = ~np.ma.getmaskarray(src)
        overlap = valid & ~np.ma.getmaskarray(dst)
        if np.any(src.data[overlap] != dst.data[overlap]):
            raise ValueError("conflicting elevations in overlapping native cells")
        dst[valid] = src[valid]
    return output, transform(Window(left, top, right - left, bottom - top), t)
