"""
One-dimensional renderers: scalars, line plots, time series and vertical profiles.

Several variables on one of these is a different problem from several on a map.
There is no colour channel to share and no per-layer colour bar to say which
scale a mark belongs to - a curve has one colour for all of it, and the only
things that can carry "which variable" are the legend and, when the units
differ, a second axis. Hence `overlay.axis_groups`, which is where the rule that
a curve plot carries two scales at most is written down.
"""

import numpy as np
import matplotlib.pyplot as plt

from .. import overlay
from ..colors import curve_color, curve_style
from ..figure import axes_for, finish_figure

# Where a secondary curve's legend entry says which axis it is read off. Text
# rather than a colour-coded spine: with two curves on the right there is no one
# colour to paint the spine with, and a wrong cue is worse than none.
SECONDARY_NOTE = {'x': ' [right axis]', 'y': ' [top axis]'}


def render_scalar(ctx):
    """
    A single value has no figure; print it.
    """
    value = float(ctx.data)
    units = f" {ctx.units}" if ctx.units else ""
    print(f"\033[36mScalar '{ctx.varname}': {value}{units}\033[0m")
    return 0


def _layers_of(ctx):
    """
    The variables this figure draws, base first.

    A plot with no `--overlay` is the same thing with one layer, so the drawing
    code below never has to ask which it is.
    """
    if ctx.composite:
        return list(ctx.composite.layers)
    return [overlay.Layer(varname=ctx.varname, data=ctx.data, colormap=ctx.colormap,
                          norm=ctx.norm, units=ctx.units, long_name=ctx.long_name)]


def _draw_layers(ax, layers, coordinate, along='x'):
    """
    Draw every layer against a shared coordinate. Returns (handles, twin).

    `along='x'` puts the coordinate on X and the values on Y - a curve against
    time. `along='y'` is the profile: the values go on X and the coordinate up
    the Y axis, so its second scale is a twin *x* axis instead.
    """
    primary, secondary, _ = overlay.axis_groups(layers)
    twin = None
    handles = []

    for index, layer in enumerate(layers):
        if layer in secondary:
            if twin is None:
                twin = ax.twinx() if along == 'x' else ax.twiny()
            target, note = twin, SECONDARY_NOTE[along]
        elif layer in primary:
            target, note = ax, ''
        else:
            continue

        values = np.asarray(layer.data, dtype=float)
        style = dict(marker='o' if values.size <= 50 else None,
                     color=curve_color(layer.requested, index),
                     linestyle=curve_style(index),
                     label=layer.label + note)
        pair = (coordinate, values) if along == 'x' else (values, coordinate)
        line, = target.plot(*pair, **style)
        handles.append(line)

    return handles, twin


def _label_axes(ax, twin, layers, coordinate_label, solo_label, along='x'):
    """
    Name each axis after what it measures.

    With one curve on an axis that is the variable and its units; with several
    it can only be the unit they share, because the legend is already naming the
    variables and repeating them here would say nothing new.
    """
    primary, secondary, _ = overlay.axis_groups(layers)

    def describe(group):
        if len(layers) == 1:
            # One curve is the plot it has always been, labelled the way it has
            # always been labelled - `Layer.label` brackets its units where the
            # axis label parenthesises them, and changing every existing figure
            # over a bracket is not an improvement.
            return solo_label
        if len(group) == 1:
            return group[0].label
        unit = overlay.units_of(group)
        return f"[{unit}]" if unit else 'value'

    if along == 'x':
        ax.set_xlabel(coordinate_label)
        ax.set_ylabel(describe(primary))
        if twin is not None:
            twin.set_ylabel(describe(secondary))
    else:
        ax.set_ylabel(coordinate_label)
        ax.set_xlabel(describe(primary))
        if twin is not None:
            twin.set_xlabel(describe(secondary))


def _keep_what_fits(layers, kind):
    """
    Drop the curves there is no axis left for, naming each and what to use
    instead.

    Per layer rather than for the whole figure: the base variable always
    survives, which is the same courtesy `--overlay` already extends to a layer
    on the wrong grid.
    """
    primary, secondary, refused = overlay.axis_groups(layers)
    for layer in refused:
        print(f"Overlay: '{layer.varname}' is in {layer.units or 'no units'} "
              f"where '{primary[0].varname}' is in "
              f"{overlay.units_of(primary) or 'no units'} and "
              f"'{secondary[0].varname}' in "
              f"{overlay.units_of(secondary) or 'no units'}; a {kind} plot "
              f"carries two scales at most, so it is not drawn. Use "
              f"--overlay-style panels for all of them.")
    return primary + secondary


def _legend(ax, handles):
    """
    One legend for both axes. A twin keeps its own handles, so they have to be
    gathered by hand or the figure comes out with two legends.

    Only when there is something to tell apart: a legend on every
    single-variable plot would be a box repeating what the axis already says,
    and would move every curve in the regression baseline for nothing.
    """
    if len(handles) > 1:
        ax.legend(handles=handles, fontsize=8, framealpha=0.85)


def _title(ctx, suffix):
    """
    The figure's title, naming every variable on it.

    `describe` says "b over a", which is stacking language: on one set of axes
    nothing is over anything, so the names go side by side instead.
    """
    if ctx.composite and not ctx.title_override:
        return f"{overlay.name_list(ctx.composite)} {suffix}".strip()
    return ctx.titled(suffix)


def _curve(ctx, x, xlabel, title):
    layers = _keep_what_fits(_layers_of(ctx), ctx.plan.kind)
    fig, ax = axes_for(ctx, (8, 4))
    handles, twin = _draw_layers(ax, layers, x, along='x')
    _label_axes(ax, twin, layers, xlabel, ctx.label, along='x')
    ax.grid(True)
    _legend(ax, handles)
    ax.set_title(title, fontweight='bold')
    return finish_figure(fig, ctx.output_path, dpi=ctx.dpi)


def render_line(ctx):
    """
    A curve against whichever coordinate remains.
    """
    axis = ctx.plan.x
    x = axis.values if axis.values is not None else np.arange(ctx.data.shape[0])
    return _curve(ctx, x, axis.label, _title(ctx, f"vs {axis.dim}"))


def render_timeseries(ctx):
    """
    A curve against time. Identical to a line plot apart from the title, but
    kept separate so the two can diverge without another conditional.
    """
    axis = ctx.plan.x
    x = axis.values if axis.values is not None else np.arange(ctx.data.shape[0])
    return _curve(ctx, x, axis.label, _title(ctx, "over time"))


def render_profile(ctx):
    """
    A vertical profile: the value on X and the vertical coordinate on Y, with
    depth increasing downward.

    The pre-refactor code drew these like any other 1D plot, with depth on X.
    """
    axis = ctx.plan.y
    y = axis.values if axis.values is not None else np.arange(ctx.data.shape[0])

    layers = _keep_what_fits(_layers_of(ctx), 'profile')
    fig, ax = axes_for(ctx, (5, 7))
    handles, twin = _draw_layers(ax, layers, y, along='y')
    _label_axes(ax, twin, layers, axis.label, ctx.label, along='y')
    ax.grid(True)
    _legend(ax, handles)
    if ctx.plan.invert_y:
        ax.invert_yaxis()
    ax.set_title(_title(ctx, 'profile'), fontweight='bold')
    return finish_figure(fig, ctx.output_path, dpi=ctx.dpi)
