"""
Working out what each dimension of a variable means.

Roles are resolved from metadata first and from names only as a last resort:

    axis attribute        confidence 100   XIOS, and anything conventions.py fixed
    CF standard_name       confidence 90   XIOS
    positive attribute     confidence 80   soildepth, altitude
    canonical units        confidence 70   everything, once normalized
    name tables            confidence 40   PEM and LMDZ files, rlonu/rlatv
    nothing                confidence 10   nslope, index, descriptor

The name tables below are the historical mechanism. They are kept because the
PEM and LMDZ writers emit no usable metadata at all, but they are now the
fallback rather than the primary source, and every Axis records which rule
resolved it so `--explain` can show the difference.
"""

from dataclasses import dataclass

import numpy as np

from .conventions import (CF_STANDARD_NAMES, LAT_NAMES, LON_NAMES, TIME_NAMES,
                          UNSTRUCTURED_DIMS)

# Constants for recognized dimension names
TIME_DIMS = ("Time", "time", "time_counter")
LAT_DIMS  = ("latitude", "lat")
LON_DIMS  = ("longitude", "lon")
# Vertical-like coordinates, kept on the Y axis of 2D cross-sections
VERT_DIMS = ("altitude", "alt", "level", "levels", "presnivs", "nb_str_max",
             "soil_layers", "subsurface_layers", "soildepth", "interlayer",
             "nlayer", "nlayer_plus_1", "ocean_layers")


def is_lat(name):
    """
    Return True if a dimension/variable name denotes latitude.
    """
    return name.lower() in LAT_DIMS


def is_lon(name):
    """
    Return True if a dimension/variable name denotes longitude.
    """
    return name.lower() in LON_DIMS


def is_time(name):
    """
    Return True if a dimension/variable name denotes time.
    """
    return name.lower() in tuple(t.lower() for t in TIME_DIMS)


def is_vertical(name):
    """
    Return True if a dimension/variable name denotes a vertical coordinate.
    """
    return name.lower() in VERT_DIMS


def wrap_longitudes(lons):
    """
    Normalize longitudes into [-180, 180] while preserving exact ±180 values.
    """
    lons = np.asarray(lons, dtype=float)
    lons = np.where(lons > 180.0, lons - 360.0, lons)
    return np.where(lons < -180.0, lons + 360.0, lons)


def find_coord_var(dataset, candidates):
    """
    Among dataset variables, return the first variable whose name matches any candidate.
    Returns None if none found.
    """
    for name in dataset.variables:
        for cand in candidates:
            if cand.lower() == name.lower():
                return name
    return None


@dataclass(frozen=True)
class Axis:
    """
    One dimension of a variable, and what it turned out to mean.

    `role` is 'X', 'Y', 'Z', 'T', 'U' (unstructured horizontal) or 'S' (some
    other index with no geometric meaning). `source` and `confidence` record
    which rule decided it.
    """
    role: str
    dim: str
    size: int
    coord: str = None
    values: np.ndarray = None
    units: str = None
    long_name: str = None
    positive: str = None
    ascending: bool = None
    is_cyclic: bool = False
    source: str = 'index'
    confidence: int = 10

    @property
    def label(self):
        """
        Axis label for a figure: the descriptive name when there is one.
        """
        base = self.long_name or self.coord or self.dim
        return f"{base} ({self.units})" if self.units else base


@dataclass(frozen=True)
class AxisRoles:
    """
    The resolved axes of one variable, in the variable's own dimension order.
    """
    axes: tuple

    def __iter__(self):
        return iter(self.axes)

    def __len__(self):
        return len(self.axes)

    def by_dim(self, dim):
        return next((a for a in self.axes if a.dim == dim), None)

    def _first(self, role):
        return next((a for a in self.axes if a.role == role), None)

    @property
    def x(self):
        return self._first('X')

    @property
    def y(self):
        return self._first('Y')

    @property
    def z(self):
        return self._first('Z')

    @property
    def t(self):
        return self._first('T')

    @property
    def u(self):
        return self._first('U')

    @property
    def others(self):
        return tuple(a for a in self.axes if a.role == 'S')

    @property
    def roles(self):
        return tuple(a.role for a in self.axes)


def _coord_for_dim(ds, dim):
    """
    The coordinate variable describing a dimension, or None.

    Only a variable named after the dimension, or one xarray recognizes as a
    coordinate, qualifies. Accepting any 1D variable on that dimension would
    make `diagevo.nc`'s `ap(interlayer)` the coordinate of its own axis, and
    every other variable on `interlayer` would then be labelled "Hybrid
    pressure".

    XIOS files carry both `time_counter` (the dimension coordinate) and
    `time_centered` (auxiliary, same dimension); the dimension coordinate wins.
    """
    if dim in ds.variables and ds[dim].dims == (dim,):
        return dim
    return next((name for name in ds.coords
                 if ds[name].dims == (dim,) and ds[name].dtype.kind in 'iufc'), None)


def _resolve_role(ds, dim, coord, unstructured):
    """
    Apply the resolution ladder to one dimension. Returns (role, source, confidence).
    """
    if unstructured and dim == unstructured.get('dim'):
        return 'U', 'unstructured', 95

    if coord is not None:
        attrs = ds[coord].attrs
        axis = attrs.get('axis')
        if axis in ('X', 'Y', 'Z', 'T'):
            return axis, 'axis-attr', 100

        sn = str(attrs.get('standard_name', ''))
        if sn in CF_STANDARD_NAMES:
            return CF_STANDARD_NAMES[sn], 'standard_name', 90

        if attrs.get('positive') in ('up', 'down'):
            return 'Z', 'positive', 80

        units = str(attrs.get('units', '')).lower()
        if units == 'degrees_north':
            return 'Y', 'units', 70
        if units == 'degrees_east':
            return 'X', 'units', 70

    # Fall back on names. The dimension name matters as much as the coordinate
    # name here: `interlayer` and `subsurface_layers` carry no coordinate
    # variable at all in the PEM files, and are still vertical axes.
    for candidate in filter(None, (coord, dim)):
        low = candidate.lower()
        if low in LAT_NAMES:
            return 'Y', 'name', 40
        if low in LON_NAMES:
            return 'X', 'name', 40
        if low in TIME_NAMES:
            return 'T', 'name', 40
        if is_vertical(low):
            return 'Z', 'name', 40
        if low in UNSTRUCTURED_DIMS:
            return 'U', 'name', 40
    return 'S', 'index', 10


def resolve_variable_axes(ds, varname, unstructured=None):
    """
    Resolve every dimension of one variable into an Axis.

    When two dimensions claim the same role, the higher confidence keeps it and
    the other is demoted to 'S', so a plot never ends up with two X axes.
    """
    da = ds[varname]
    resolved = []
    for dim, size in zip(da.dims, da.shape):
        coord = _coord_for_dim(ds, dim)
        role, source, confidence = _resolve_role(ds, dim, coord, unstructured)

        values = units = long_name = positive = None
        ascending = None
        cyclic = False
        if coord is not None:
            cvar = ds[coord]
            values = np.asarray(cvar.values, dtype=float) if cvar.dtype.kind in 'iufc' else None
            units = cvar.attrs.get('units')
            long_name = cvar.attrs.get('long_name') or cvar.attrs.get('title')
            positive = cvar.attrs.get('positive')
            if values is not None and values.size > 1:
                diffs = np.diff(values)
                if np.all(diffs > 0):
                    ascending = True
                elif np.all(diffs < 0):
                    ascending = False
                if role == 'X':
                    cyclic = looks_cyclic(values)

        resolved.append(Axis(role=role, dim=dim, size=size, coord=coord, values=values,
                             units=units, long_name=long_name, positive=positive,
                             ascending=ascending, is_cyclic=cyclic,
                             source=source, confidence=confidence))

    return AxisRoles(axes=tuple(_break_ties(resolved)))


def _break_ties(axes):
    """
    Keep at most one axis per geometric role, the most confidently resolved one.
    """
    best = {}
    for axis in axes:
        if axis.role == 'S':
            continue
        current = best.get(axis.role)
        if current is None or axis.confidence > current.confidence:
            best[axis.role] = axis

    out = []
    for axis in axes:
        if axis.role != 'S' and best[axis.role] is not axis:
            out.append(Axis(role='S', dim=axis.dim, size=axis.size, coord=axis.coord,
                            values=axis.values, units=axis.units,
                            long_name=axis.long_name, positive=axis.positive,
                            ascending=axis.ascending, source='demoted',
                            confidence=axis.confidence))
        else:
            out.append(axis)
    return out


def looks_cyclic(values):
    """
    Whether a longitude axis wraps the globe, so a map can close its seam.
    """
    values = np.asarray(values, dtype=float)
    if values.size < 2:
        return False
    step = np.median(np.diff(values))
    if step <= 0:
        return False
    span = values[-1] - values[0] + step
    return abs(span - 360.0) < 1.5 * abs(step)


def needs_seam_column(lons):
    """
    Whether a sorted longitude axis goes right round but stops short of closing.

    "Right round" is `looks_cyclic`, the same judgement a map makes from the
    resolved axis, so a grid cannot be cyclic on a map and not on the globe.
    What is left to decide is only whether the seam is still open: an axis
    holding both -180 and +180 has closed itself already, and repeating its
    first column would make a zero-width cell, which an interpolator refuses
    outright and a mesher turns into a degenerate face.
    """
    lons = np.asarray(lons, dtype=float)
    return bool(looks_cyclic(lons)) and (lons[-1] - lons[0]) < 360.0 - 1e-6


def close_longitude_seam(lons, values, axis=-1):
    """
    (lons, values) with the first column repeated a full turn later, when the
    axis goes right round but stops short of closing. Otherwise unchanged.

    Everything drawn on a sphere from a global grid needs this, and each of the
    three places that draws one found out the hard way: without it the cells
    between the last longitude and the first are off the end of the axis. A
    model grid of 32 longitudes ends at 168.75, so an interpolator returned an
    eleven degree wedge of missing data - which on a shell is an opaque red
    slice - and a mesher simply left the wedge out, cutting every cloud that
    crosses the date line into two with a flat face on each side.
    """
    lons = np.asarray(lons, dtype=float)
    values = np.asarray(values)
    if not needs_seam_column(lons):
        return lons, values
    first = np.take(values, [0], axis=axis)
    return np.append(lons, lons[0] + 360.0), np.concatenate([values, first], axis=axis)


def centers_to_edges(arr):
    """
    Convert 1D monotonic center coordinates to edges.
    Returns the edges (always ascending) and whether the input was ascending.
    """
    arr = np.asarray(arr, dtype=float)
    if arr.size < 2:
        raise ValueError("Need at least 2 coordinates to infer pcolormesh cells")

    diffs = np.diff(arr)
    if np.all(diffs > 0):
        ascending = True
    elif np.all(diffs < 0):
        ascending = False
    else:
        raise ValueError("Coordinate array must be monotonic for hover mapping")

    if not ascending:
        arr = arr[::-1]
    edges = np.empty(arr.size + 1, dtype=float)
    edges[1:-1] = 0.5 * (arr[:-1] + arr[1:])
    edges[0] = arr[0] - 0.5 * (arr[1] - arr[0])
    edges[-1] = arr[-1] + 0.5 * (arr[-1] - arr[-2])
    return edges, ascending
