"""
Choosing a colormap and a normalization from the data.

The historical default was `jet`, which is not perceptually uniform: it invents
banding where the data is smooth and hides detail at the ends. The default here
is a perceptual sequential map for ordinary fields and a symmetric diverging map
for fields that genuinely straddle zero, so that the colour of a cell means the
same thing everywhere in the figure.

`-c <name>` still passes any matplotlib colormap through untouched, so `-c jet`
reproduces the old figures exactly.
"""

import re

import numpy as np
import matplotlib.pyplot as plt
from matplotlib.colors import (CenteredNorm, LinearSegmentedColormap, LogNorm,
                               Normalize, SymLogNorm, to_hex)

SEQUENTIAL = 'viridis'
DIVERGING = 'RdBu_r'

# One colormap per overlaid variable, in the order they are given.
#
# Each runs from pale to saturated in a single hue, which is what a stack of
# overlays needs and what `viridis` cannot give: a perceptual map sweeps through
# several hues, so two of them overlaid produce colours belonging to both scales
# and to neither. A single-hue ramp keeps "which variable" in the hue and "how
# much" in the depth, and the two stay separable where the layers cross.
LAYER_COLORMAPS = ('Blues', 'Reds', 'YlOrBr', 'Greens', 'Purples', 'bone')

# One colour per layer, for everything that is a line rather than a fill: a
# curve on a 1-D plot, a set of contours over a map, the edge of a hatch.
#
# A colormap is the wrong object for any of these. A ramp needs a scale to
# spread itself over, and a line has one colour for all of it. These are the
# Okabe-Ito set, which stays separable under the common colour-vision
# deficiencies - the same reason `jet` is not the default here.
CURVE_COLORS = ('#0072B2', '#D55E00', '#009E73', '#CC79A7',
                '#E69F00', '#56B4E9', '#000000')

# Once the colours run out, the dash pattern turns over instead, so a printed
# figure in one ink still separates them.
CURVE_STYLES = ('-', '--', '-.', ':')

# The part of a colormap an atmospheric shell is drawn in, and the part of a
# grey ramp the terrain under it gets.
#
# A shell coloured over the whole range puts its outermost surface - the one
# actually in view - at the bottom of the scale, which on `viridis` is a dark
# purple that is all but invisible against the globe. Keeping to the top half
# leaves the outer envelope a legible teal and the core still yellow.
#
# The terrain is the other half of the same problem. It used to be drawn in
# `bone`, which runs black to white through blue - exactly where viridis's dark
# end lives - so shell and ground came out the same colour. A neutral grey,
# stopping short of both black and white, has no hue to be confused with and
# leaves the bright end of the picture to the shell.
SHELL_RANGE = (0.45, 1.0)
TERRAIN_RANGE = (0.15, 0.72)

# Where a colormap is sampled when one has to become a single line colour: near
# the saturated end, where the hue is unambiguous, but short of the very top,
# which on several ramps is almost white.
SATURATED_POINT = 0.75

# Names that suggest a signed quantity: PEM tendencies (d_h2oice, d_co2ice,
# d_precip, d_evap, d_runoff), differences and anomalies.
SIGNED_NAME = re.compile(r'(^d_|_diff$|tendency|anomaly|_anom|delta)', re.IGNORECASE)

# A field must have this fraction of its range on each side of zero before it is
# treated as diverging, so that a positive field with a few negative rounding
# artefacts is not centred on zero.
SIGNED_BALANCE = 0.05

# How much of a field a log scale has to be able to place before one is chosen
# automatically. A log scale has nowhere to put a zero, so every zero cell comes
# out blank: below this fraction the picture is mostly holes, and what it shows
# is where the variable is rather than what it is. `co2ice` is zero over 97% of
# a PEM step and was being drawn on a log scale that placed 33 cells of 1089.
#
# Half is the line because the claim "this field spans several decades" has to
# be about the field, not about the minority of it that happens to be nonzero.
# An explicit --norm log is still honoured - it reports what it cannot draw
# rather than second-guessing the request.
LOG_COVERAGE = 0.5

# The decades a log colour bar can be read over, and how much of a field has to
# lie inside them before one is chosen.
#
# The percentile test in `spans_decades` sees how far apart the tails are, and a
# field with a numerical floor under it makes that ratio say anything at all:
# `h2o_ice` in start.nc bottoms out at 1e-40 - what is left of a mixing ratio
# after a model runs out of precision, not a value anybody plotted - and its 5th
# and 95th percentiles are 34 decades apart. The log scale that followed spent
# two thirds of its bar on that dust and squeezed the ice into the top third.
#
# So a field has to occupy the decades, not merely reach across them. Six is the
# same budget the globe's column integral already works to
# (`render.globe.shell.MAX_COLUMN_DECADES`), and half is the line for the same
# reason as LOG_COVERAGE above: the claim "this field spans several decades" has
# to be about the field, not about the minority of it that happens to be spread
# out. `h2o_ice` has 25% of its cells inside six decades of its 95th percentile;
# `soildepth`, whose thin top layers are real and not dust, has 79%.
LOG_DECADES = 6
LOG_DECADE_SHARE = 0.5

# Above this fraction of non-positive cells the automatic log scale becomes a
# symlog instead, which has somewhere to put a zero. Below it, a handful of stray
# zeros is not worth the extra structure a symlog puts in the colour bar, and
# those few cells are painted rather than placed.
SYMLOG_SWITCH = 0.02

# What a cell a log scale cannot place is drawn in. Left to itself matplotlib
# draws nothing there, and a hole in a map reads as "the field is not here" -
# a different claim from "the field is zero here". A grey that belongs to no
# colormap says "off the scale", which is the true one.
BLANK_COLOR = '0.82'


def truncated(colormap, low, high, steps=256):
    """
    A colormap using only the `low`..`high` part of another's range.

    Used where two colour scales share a picture and have to stay tellable
    apart - the globe's terrain under an atmospheric shell, and the shell over
    it. Restricting the range is not the same as rescaling the data: the values
    still span the whole scale and the bar still reads in real units, only the
    colours the scale is drawn in are narrowed.

    Takes a name or a Colormap, because `--cmap` passes through whatever was
    typed.
    """
    base = plt.get_cmap(colormap) if isinstance(colormap, str) else colormap
    name = f"{getattr(base, 'name', 'cmap')}_{low:g}_{high:g}"
    return LinearSegmentedColormap.from_list(
        name, base(np.linspace(low, high, steps)))


def curve_color(requested, index):
    """
    One colour for one layer drawn as lines rather than as a fill.

    `--overlay VAR:NAME` names a colormap on a blended map; on a curve or a
    contour set the same word has to become a single colour. A real colour name
    is taken as itself, so `--overlay co2_ice:red` does what it looks like; a
    colormap is sampled near its saturated end, so `--overlay co2_ice:Reds`
    still draws a red line and a variable keeps the same hue whichever style it
    is drawn in. With nothing named, the rota answers.
    """
    if requested:
        try:
            return to_hex(requested)
        except ValueError:
            pass
        try:
            return to_hex(plt.get_cmap(requested)(SATURATED_POINT))
        except (ValueError, KeyError):
            print(f"Warning: '{requested}' is neither a colour nor a colormap; "
                  f"using the next colour in the rota.")
    return CURVE_COLORS[index % len(CURVE_COLORS)]


def curve_style(index):
    """
    The dash pattern for one layer, turning over only once the colours have.
    """
    return CURVE_STYLES[(index // len(CURVE_COLORS)) % len(CURVE_STYLES)]


def _finite(data):
    values = np.asarray(data, dtype=float).ravel()
    return values[np.isfinite(values)]


def straddles_zero(data, balance=SIGNED_BALANCE):
    """
    Whether the data meaningfully spans both signs.
    """
    values = _finite(data)
    if values.size == 0:
        return False
    low, high = values.min(), values.max()
    if low >= 0 or high <= 0:
        return False
    span = high - low
    return span > 0 and (-low / span) >= balance and (high / span) >= balance


def symmetric_limit(data):
    """
    A robust symmetric limit for a diverging scale: the larger of the 1st and
    99th percentile magnitudes, so a single outlier cannot flatten the figure.
    """
    values = _finite(data)
    if values.size == 0:
        return 1.0
    low, high = np.percentile(values, [1, 99])
    limit = max(abs(low), abs(high))
    return float(limit) if limit > 0 else float(max(abs(values).max(), 1e-12))


def spans_decades(data, decades=3):
    """
    Whether the bulk of a positive field covers enough orders of magnitude to be
    worth a log scale.

    The test deliberately uses the 5th and 95th percentiles rather than the
    minimum and maximum. Plenty of perfectly linear fields touch zero somewhere
    - a wind speed vanishes at the poles, an ice thickness over bare ground -
    and keying on the smallest positive value would put those on a log scale
    where the whole map collapses to one colour.

    A field must also be mostly *there*. Percentiles of the positive values say
    nothing about how many of them there are, so a field that is zero almost
    everywhere and spans decades across the handful of cells that are not - a
    CO2 ice cap, which is 3% of a PEM step - passed this test and was drawn on a
    log scale that left 97% of the map blank. See LOG_COVERAGE.

    And the decades have to be real. A ratio of two percentiles says how far
    apart the tails are, which a numerical floor under the field inflates
    without bound: `h2o_ice` in start.nc reaches 1e-40, so its percentiles are
    34 decades apart and the bar that followed spent two thirds of itself on
    precision loss. The field must occupy the decades it reaches across. See
    LOG_DECADES.
    """
    values = _finite(data)
    if values.size < 2 or np.any(values < 0):
        return False
    positive = values[values > 0]
    if positive.size < 2 or positive.size < LOG_COVERAGE * values.size:
        return False
    low, high = np.percentile(positive, [5, 95])
    if low <= 0 or high / low < 10 ** decades:
        return False
    return (positive >= high / 10.0 ** LOG_DECADES).mean() >= LOG_DECADE_SHARE


def choose_style(data, varname, cmap='auto', norm=None, vmin=None, vmax=None,
                 signed=False):
    """
    Return (colormap, matplotlib_norm, explanation).

    The colormap is the name that was chosen, except where a log scale has cells
    it cannot place: those come back as a Colormap object marked with how many,
    which is what tells the renderer to lay a colour under them instead of
    leaving holes. See `_with_blank_color` and `log_colormap`.

    `signed` marks data that is a difference by construction, such as the output
    of --anomaly, which is centred on zero regardless of what the values do.

    One exit, so that the blank-cell colour is applied to every branch that can
    end up on a log scale - the automatic one and `--norm log` alike - rather
    than being repeated in each of them and forgotten in one.
    """
    name, chosen, why = _style(data, varname, cmap, norm, vmin, vmax, signed)
    return _with_blank_color(name, chosen, data), chosen, why


def _style(data, varname, cmap, norm, vmin, vmax, signed):
    """
    The choice itself: (colormap_name, norm, explanation).
    """
    explicit_limits = vmin is not None or vmax is not None

    # A difference field is centred on zero whatever palette was asked for.
    # `-c` chooses the colours; it does not get to say what the middle colour
    # means, and a diverging map whose white sits at some arbitrary fraction of
    # the bar is worse than no diverging map at all. `signed` is set only by
    # --anomaly and --diff, which are explicit requests for a difference, so
    # this overrides an explicit colormap without ever overriding a guess.
    if signed:
        limit = symmetric_limit(data)
        chosen = _explicit_norm(data, norm, vmin, vmax) if explicit_limits or norm \
            else Normalize(vmin=-limit, vmax=limit)
        name = cmap if cmap and cmap != 'auto' else DIVERGING
        return name, chosen, "difference field, centred on zero"

    # An explicit colormap is honoured exactly, with no automatic norm games.
    # The name heuristic below stays under this line deliberately: that one is a
    # guess about what a variable means from what it is called, and a guess must
    # not override what the user typed.
    if cmap and cmap != 'auto':
        return cmap, _explicit_norm(data, norm, vmin, vmax), f"colormap '{cmap}' requested"

    if SIGNED_NAME.search(varname or '') and straddles_zero(data):
        limit = symmetric_limit(data)
        chosen = _explicit_norm(data, norm, vmin, vmax) if explicit_limits or norm \
            else Normalize(vmin=-limit, vmax=limit)
        return DIVERGING, chosen, f"'{varname}' looks signed and spans both signs"

    if norm is None and not explicit_limits and spans_decades(data):
        values = _finite(data)
        positive = values[values > 0]
        blanked = values.size - positive.size

        # Nobody asked for a log scale here - this branch decided on one because
        # the field spans decades. Where that decision would have cost a real
        # part of the map, it takes the scale that can draw the whole of it
        # instead: symlog is log over the decades and linear across the window
        # under the smallest positive value, so a zero lands at the bottom of
        # the ramp rather than nowhere. An explicit --norm log is never
        # rewritten this way; see `_explicit_norm`.
        if blanked > SYMLOG_SWITCH * values.size:
            return SEQUENTIAL, _log_like_norm(values, positive), \
                "spans several decades and contains zeros, symlog scale"

        _report_blanked_by_log(values, positive)
        return SEQUENTIAL, LogNorm(vmin=positive.min(), vmax=positive.max()), \
            "positive field spanning several decades, log scale"

    return SEQUENTIAL, _explicit_norm(data, norm, vmin, vmax), "sequential default"


def _log_like_norm(values, positive):
    """
    A log scale with a floor under it, for a field whose zeros matter.

    Asymmetric, unlike the `--norm symlog` branch below: that one answers a
    request about a signed field and is built from `symmetric_limit`, while this
    one is reached only for a field with no negatives at all, where a symmetric
    scale would spend half the colour bar on values that do not exist.

    `vmin` is zero and `linthresh` sits at the smallest positive value, so the
    linear window holds nothing but the zeros themselves and the decades above
    it are exactly the ones the LogNorm would have covered. A wider window would
    compress the small values into the bottom of the ramp on the way to a
    prettier bar, which is a different picture of the field, not the same one
    drawn better.
    """
    high = float(positive.max())
    low = float(min(values.min(), 0.0))
    linthresh = float(positive.min())
    if not np.isfinite(linthresh) or linthresh <= 0:
        linthresh = max(abs(high) * 1e-6, 1e-12)
    return SymLogNorm(linthresh=linthresh, vmin=low, vmax=high)


def blank_cells(values):
    """
    Where a log scale has nothing to place: the finite, non-positive cells.

    Finite deliberately. A NaN is missing data, and a hole is exactly what
    missing data should look like; it is the zeros that need a colour, and
    telling the two apart is the whole point of giving them one.
    """
    array = np.asarray(values, dtype=float)
    return np.isfinite(array) & (array <= 0)


def log_colormap(name, blanked, total):
    """
    A copy of a colormap, marked as having cells a log scale cannot place.

    The mark travels on the colormap because that is the object still in hand
    everywhere it is needed: `render.layers.paint_blank_cells` reads it off
    before drawing the mesh and lays the grey under it, and
    `figure.annotate_blank_cells` reads it off afterwards and writes the note
    below the picture. A copy, because the alternative is marking the colormap
    matplotlib hands out to everything else as well.

    Not `set_bad`, which is the obvious way and the wrong one: cartopy needs a
    fully transparent 'bad' to wrap a quadmesh across the dateline, and when it
    does not get one it warns and redraws the wrapped copy over the entire map.
    """
    base = plt.get_cmap(name) if isinstance(name, str) else name
    cmap = base.copy()
    cmap._dispnc_blank_note = (
        f"grey: {blanked} of {total} cells ({blanked / total:.0%}) at zero or "
        f"below, off the log scale")
    return cmap


def _with_blank_color(name, norm, data):
    """
    Give a log scale's undrawable cells a colour of their own, or pass the name
    through untouched.

    Only a LogNorm has this problem - it is the one scale here with nowhere to
    put a zero - and only when the field actually contains one, so an ordinary
    figure still gets a plain colormap name and the transparent 'bad' colour
    matplotlib gives it, which is what a genuine gap in the data should look
    like.
    """
    if not isinstance(norm, LogNorm):
        return name
    values = _finite(data)
    if not values.size:
        return name
    blanked = int(values.size - (values > 0).sum())
    return log_colormap(name, blanked, values.size) if blanked else name


def _report_blanked_by_log(values, positive):
    """
    Say how much of a field a log scale had to put outside its own scale.

    A log scale has nowhere to put a zero or a negative, so those cells are
    drawn in `BLANK_COLOR` rather than in a colour off the bar - and a tendency,
    which is exactly what gets reached for when values span decades, is negative
    over a third of a typical map. `d_h2oice` on one PEM step is 335 cells of
    1089. The figure says so too, under the picture; this says it where a batch
    run will see it.

    The automatic choice avoids the situation altogether once it is worth
    avoiding (see SYMLOG_SWITCH). An explicit --norm log is still never rewritten
    into a symlog: this line is the whole of the answer there, because silently
    changing the scale somebody named is the quiet substitution this module
    refuses everywhere else.
    """
    blanked = values.size - positive.size
    if not blanked or not values.size:
        return
    print(f"Note: a log colour scale cannot place {blanked} of {values.size} cells "
          f"({blanked / values.size:.0%}), which are zero or negative; they are "
          f"drawn in grey. Use --norm symlog to put them on the scale.")


def _explicit_norm(data, norm, vmin, vmax):
    """
    Build the normalization the user asked for, if any.
    """
    if norm in (None, 'linear'):
        return Normalize(vmin=vmin, vmax=vmax) if (vmin is not None or vmax is not None) else None
    if norm == 'log':
        values = _finite(data)
        positive = values[values > 0]
        if positive.size == 0:
            print("Note: --norm log needs positive values; falling back to linear.")
            return None
        _report_blanked_by_log(values, positive)
        return LogNorm(vmin=vmin if vmin is not None else positive.min(),
                       vmax=vmax if vmax is not None else positive.max())
    if norm == 'symlog':
        limit = symmetric_limit(data)
        return SymLogNorm(linthresh=max(limit * 1e-3, 1e-12),
                          vmin=vmin if vmin is not None else -limit,
                          vmax=vmax if vmax is not None else limit)
    if norm == 'centered':
        limit = symmetric_limit(data)
        return Normalize(vmin=vmin if vmin is not None else -limit,
                         vmax=vmax if vmax is not None else limit)
    return None
