"""
Summary statistics for a field.

The global mean is area-weighted wherever the geometry allows it. An unweighted
mean over a latitude/longitude grid is dominated by the poles, where the cells
are tiny, and is simply the wrong number for a planetary average.
"""

import numpy as np

PERCENTILES = (1, 5, 25, 50, 75, 95, 99)


def area_weights(ds, report, axes, shape, reduced=()):
    """
    Weights matching `shape`, or (None, reason) when the field is not spatial.

    Prefers the file's own cell-area variable, which exists as `cell_area` in
    PEM files and `aire`/`area` in LMDZ files, and falls back to cos(latitude)
    for XIOS output, which ships no area variable at all.

    `reduced` names the dimensions a `--reduce` has already collapsed. A zonal
    mean is still a field over latitude, and weighting it still matters: the
    unweighted average of a latitude profile is dominated by the poles in
    exactly the way this module exists to prevent. Without this the note read
    "no latitude axis" and the number quietly changed meaning.
    """
    lat_axis = next((a for a in axes if a.role == 'Y'), None)

    if report.cell_area and report.cell_area in ds.variables:
        area = np.asarray(ds[report.cell_area].values, dtype=float)
        if area.shape == shape:
            return area, f"cell_area variable '{report.cell_area}'"
        if area.T.shape == shape:
            return area.T, f"cell_area variable '{report.cell_area}' (transposed)"

    if lat_axis is not None and lat_axis.values is not None:
        weights = np.cos(np.deg2rad(np.asarray(lat_axis.values, dtype=float)))
        weights = np.clip(weights, 0.0, None)
        if len(shape) == 2:
            if shape[0] == weights.size:
                return np.repeat(weights[:, None], shape[1], axis=1), "cos(latitude)"
            if shape[1] == weights.size:
                return np.repeat(weights[None, :], shape[0], axis=0), "cos(latitude)"
        # What a reduction leaves: a profile down the latitude axis, which the
        # same cosine still weights correctly.
        if len(shape) == 1 and shape[0] == weights.size:
            return weights, "cos(latitude)"

    if reduced:
        return None, (f"unweighted (the weights do not survive reducing over "
                      f"{', '.join(reduced)})")
    return None, "unweighted (no area variable and no latitude axis)"


def summarize(data, weights=None):
    """
    Descriptive statistics, ignoring NaN throughout.
    """
    values = np.asarray(data, dtype=float)
    flat = values.ravel()
    finite = flat[np.isfinite(flat)]

    out = {
        'count': int(flat.size),
        'valid': int(finite.size),
        'missing': int(flat.size - finite.size),
    }
    if finite.size == 0:
        return out

    out.update({
        'min': float(finite.min()),
        'max': float(finite.max()),
        'mean': float(finite.mean()),
        'std': float(finite.std()),
    })
    for p, value in zip(PERCENTILES, np.percentile(finite, PERCENTILES)):
        out[f'p{p}'] = float(value)

    if weights is not None and weights.shape == values.shape:
        good = np.isfinite(values)
        # The weights of missing cells must be excluded from the denominator as
        # well as the numerator, which is the classic way to get this wrong.
        denominator = float(np.sum(np.where(good, weights, 0.0)))
        if denominator > 0:
            numerator = float(np.sum(np.where(good, values * weights, 0.0)))
            out['weighted_mean'] = numerator / denominator
    return out


def print_summary(varname, data, weights=None, weight_note='', units=None):
    """
    Print the summary in a fixed-width block.
    """
    stats = summarize(data, weights)
    suffix = f" [{units}]" if units else ""
    print(f"\nStatistics for '{varname}'{suffix}")
    print(f"  points        : {stats['count']} ({stats['valid']} valid, "
          f"{stats['missing']} missing)")
    if stats['valid'] == 0:
        print("  no finite values")
        return stats
    print(f"  min / max     : {stats['min']:.6g} / {stats['max']:.6g}")
    print(f"  mean / std    : {stats['mean']:.6g} / {stats['std']:.6g}")
    if 'weighted_mean' in stats:
        print(f"  area-weighted : {stats['weighted_mean']:.6g}   [{weight_note}]")
    else:
        print(f"  area-weighted : not available   [{weight_note}]")
    percentiles = '  '.join(f"p{p}={stats[f'p{p}']:.4g}" for p in PERCENTILES)
    print(f"  percentiles   : {percentiles}")
    return stats
