"""
Latitude/longitude maps.

These are drawn on a cartopy GeoAxes rather than plain axes. The pre-refactor
code used plain axes, which forced three separate workarounds: the topography
had to be redrawn in [0, 360] whenever a file used that convention, the overlay
was skipped entirely when latitude landed on X, and the seam between the last
and first longitude column stayed open. Putting the data in a real projection
removes all three: longitudes are normalized once, and everything else is drawn
through `transform=PlateCarree()`.
"""

import numpy as np
import matplotlib.pyplot as plt
import cartopy.crs as ccrs
from cartopy.util import add_cyclic_point

from .. import overlay
from ..coords import needs_seam_column, wrap_longitudes
from ..figure import axes_for, finish_figure, attach_format_coord, is_video
from ..topography import overlay_topography
from . import extras, layers as layer_draw
from .movie import drive, frame_label


# Where a map's title sits, in figure coordinates.
TITLE_Y = 0.93

# Where a panel's own title sits, in that panel's axes coordinates. Above the
# top gridline labels, which are turned off, but clear of the frame.
PANEL_TITLE_Y = 1.04


def _normalize_longitudes(lons, data):
    """
    Put longitudes on [-180, 180] ascending, reordering the data columns to
    match. Files using [0, 360] and files using [-180, 180] then look the same
    to everything downstream, including the topography overlay.
    """
    lons = wrap_longitudes(np.asarray(lons, dtype=float))
    order = np.argsort(lons, kind='stable')
    return lons[order], data[..., order]


def _equally_spaced(lons, tolerance=1e-3):
    """
    Whether the longitude steps are uniform.

    Wrapping a staggered grid can break uniformity: start.nc's `rlonu` runs past
    180 degrees, so normalizing it moves its last points to the front and leaves
    one irregular step. add_cyclic_point rejects such an axis, so it must not be
    offered one.
    """
    if lons.size < 3:
        return False
    steps = np.diff(lons)
    return bool(np.all(np.abs(steps - steps[0]) <= tolerance * max(abs(steps[0]), 1e-12)))


def render_geomap(ctx):
    """
    Draw the field on a PlateCarree map, with optional topography contours, and
    then offer the polar and globe views.
    """
    plan = ctx.plan
    data = ctx.data

    source_lons = plan.x.values
    lats = plan.y.values
    if source_lons is None:
        source_lons = np.arange(data.shape[1], dtype=float)
    if lats is None:
        lats = np.arange(data.shape[0], dtype=float)

    # Everything drawn on these axes goes through the same longitude handling:
    # the base field, its animation frames, and every overlay layer. `_prepare`
    # is the single place that happens, because the reordering is not always the
    # identity - start.nc's `rlonu` runs to 185.6 degrees, so wrapping moves its
    # last column to the front - and anything that skipped it would be drawn
    # against longitudes it does not belong to. Two identical fields came out a
    # column apart. Both helpers act on the last axis, so a leading frame axis
    # passes through untouched.
    frames = ctx.frames
    lons, data = _normalize_longitudes(source_lons, data)
    if frames is not None:
        _, frames = _normalize_longitudes(source_lons, frames)

    # Closes the gap between the last and first column, which is 11.25 degrees
    # wide on the XIOS grids - but only where there is still a gap. An axis
    # holding both -180 and +180 has closed itself already, and `needs_seam_column`
    # is the same question the globe asks before repeating a column. Without it
    # `add_cyclic_point` appends a 34th column at 191.25 degrees on a 33-column
    # file that already went right round: a duplicate of the first column, laid
    # over the seam and hanging past the dateline. Per-cell that column was
    # quietly redrawn by cartopy's wrap handling and looked like nothing;
    # interpolated it is a visible seam at 180.
    cyclic = bool(plan.cyclic and _equally_spaced(lons) and needs_seam_column(lons))
    if cyclic:
        data, lons = add_cyclic_point(data, coord=lons)
        if frames is not None:
            frames = add_cyclic_point(frames)

    layers = _prepare_layers(ctx.composite, source_lons, cyclic)

    proj = ccrs.PlateCarree()
    # Wider per colour bar that will actually be drawn, so they do not run off
    # the page. A contour or a hatch layer needs none, which is part of what
    # those styles buy: the width goes back to the map.
    extra = layer_draw.bar_count(ctx.composite, ctx.overlay_style) - 1
    fig, ax = axes_for(ctx, (8 + 1.1 * extra, 6), projection=proj)
    if ctx.panel is None:
        # The bars are insets, so they are measured in fractions of the axes and
        # would run off the canvas if the axes kept the full default width. In a
        # panel the grid owns the spacing and this would fight it.
        fig.subplots_adjust(right=0.86 - 0.07 * extra)
    mesh, wants_bar = layer_draw.draw_base(
        fig, ax, data, layers, lons, lats, ctx.colormap, ctx.norm, ctx.label,
        style=ctx.overlay_style, transform=proj, interpolate=ctx.interpolate)

    drawn = layer_draw.draw(fig, ax, layers, lons, lats,
                            style=ctx.overlay_style, transform=proj,
                            interpolate=ctx.interpolate)

    try:
        attach_format_coord(ax, data, lons, lats, 'lon', 'lat', ctx.varname)
    except ValueError:
        pass

    if ctx.show_topo:
        overlay_topography(ax, transform=proj, levels=10)

    gl = ax.gridlines(draw_labels=True, linewidth=0.4, color='gray',
                      alpha=0.6, linestyle='--')
    gl.top_labels = False
    gl.right_labels = False

    # An inset bar is measured against its own axes and reaches past the right
    # of it, which over a lone map is empty canvas and inside a grid is the next
    # panel. There, a bar that takes its room from the panel is the only one
    # that stays in it.
    if not wants_bar:
        pass          # a bivariate scheme has its square key instead
    elif ctx.panel is None:
        fig.colorbar(mesh, cax=layer_draw.bar_axes(ax, 0)).set_label(ctx.label)
    else:
        fig.colorbar(mesh, ax=ax, pad=0.05, shrink=0.85).set_label(ctx.label)
    _set_title(ctx, fig, ax, f"{ctx.title}\n{overlay.describe(ctx.composite)}"
               if drawn else ctx.title)

    consumed = False
    if frames is not None:
        consumed = _animate(fig, ax, mesh, ctx, frames, lons, lats, drawn)

    # Save the main map before opening any secondary view. A failure to write it
    # is carried to the end rather than returned here, so that --show-polar and
    # --show-3d still get their chance: they write their own files, and one
    # unwritable path should not silently cost the others.
    status = 0
    if not consumed:
        status = finish_figure(fig, ctx.output_path, dpi=ctx.dpi)

    lon2d, lat2d = np.meshgrid(lons, lats)
    extras.show_extra_views(lon2d, lat2d, data, ctx.colormap, ctx.varname, ctx.units,
                            ctx.interactive, ctx.show_polar, ctx.show_3d,
                            ctx.show_topo, ctx.output_path, norm=ctx.norm,
                            globe_options=ctx.globe,
                            frames=frames, frame_dim=ctx.frame_dim,
                            frame_axis=ctx.frame_axis, fps=ctx.fps,
                            composite=ctx.composite, layers=layers,
                            overlay_style=ctx.overlay_style,
                            interpolate=ctx.interpolate)
    return status


def _set_title(ctx, fig, ax, text):
    """
    Title a map.

    On the figure rather than on the axes, because a cartopy GeoAxes here
    resolves its title's position to NaN and the text is then never rasterized -
    present, visible, and invisible. The figure's own title uses an ordinary
    transform and lands where it should.

    Except in a panel, where the figure's title belongs to the grid and every
    panel writing to it would leave only the last one's. There the title has to
    go above its own axes - but not via `set_title`, which is the very call that
    resolves to NaN: measured, `gridlines(draw_labels=True)` is what does it, and
    a map without its degree labels is not the trade. `ax.text` against
    `transAxes` is an ordinary transform, so it lands like the figure title does.
    """
    if ctx.panel is not None:
        ax.text(0.5, PANEL_TITLE_Y, text, transform=ax.transAxes,
                ha='center', va='bottom', fontweight='bold', fontsize=10)
    else:
        fig.suptitle(text, fontweight='bold', y=TITLE_Y)



def _prepare_layers(composite, source_lons, cyclic):
    """
    Every `--overlay` layer put through the longitude handling the base field
    has just been through. Returns [(layer, values, frames)].

    A layer arrives on the variable's own axis order, exactly as the base did,
    so it needs exactly the same wrap-and-reorder and the same cyclic column.
    Doing it here rather than at draw time is what keeps the animation honest
    too: `_animate` swaps in `frames[index]`, which must already be in the drawn
    order or the movie would step through mis-registered fields.
    """
    prepared = []
    for layer in (composite.layers[1:] if composite else ()):
        _, values = _normalize_longitudes(source_lons,
                                          np.asarray(layer.data, dtype=float))
        frames = None
        if layer.frames is not None:
            _, frames = _normalize_longitudes(source_lons,
                                              np.asarray(layer.frames, dtype=float))
        if cyclic:
            values = add_cyclic_point(values)
            if frames is not None:
                frames = add_cyclic_point(frames)
        prepared.append((layer, values, frames))
    return prepared


def _animate(fig, ax, mesh, ctx, frames, lons, lats, drawn=()):
    """
    Step the map through `ctx.frames`: a movie with `-o movie.mp4`, a slider
    otherwise. Returns True when the figure was written and closed here.

    Only the meshes' arrays, the title and the hover readout change; the
    projection, the topography overlay and the colour scales are the same in
    every frame, which is what makes a 669-step year cheap to render.

    Every overlay steps with the base. They used to be drawn once and left
    there, so a year of surface pressure ran underneath an ice field frozen on
    step zero - a movie that was wrong in a way nothing announced.
    """
    # The hover readout exists for a cursor, and a movie has none. Rebuilding it
    # per frame is the most expensive thing in the loop after the draw itself.
    scrubbable = not is_video(ctx.output_path)

    def on_frame(index):
        mesh.set_array(frames[index])
        layer_draw.refill_blank_cells(mesh, frames[index])
        for layer, artist, layer_frames in drawn:
            if layer_frames is None or artist is None:
                # No frames means a layer that carries no animated dimension - a
                # mask, a topography - which keeps what it was drawn with rather
                # than vanishing when the field under it starts moving. No artist
                # means a style whose geometry cannot be refilled in place; those
                # hold their first frame too, and say so below.
                continue
            values = layer_frames[index]
            artist.set_array(values)
            # The transparency is derived from the values, so leaving it behind
            # would show this frame's field through the last one's holes
            artist.set_alpha(overlay.layer_alpha(values, layer.norm))
        _set_title(ctx, fig, ax, f"{ctx.title}\n"
                   f"{frame_label(ctx.frame_axis, ctx.frame_dim, index)}")
        # The readout closes over the array it was given, so scrubbing the
        # slider would otherwise keep reporting the first frame's values under
        # a map showing the last one
        if scrubbable:
            try:
                attach_format_coord(ax, frames[index], lons, lats, 'lon', 'lat', ctx.varname)
            except ValueError:
                pass
        return (mesh, *(a for _, a, _ in drawn if a is not None))

    on_frame(0)
    return drive(fig, ctx, on_frame, ctx.output_path,
                 label=ctx.frame_dim or 'frame')
