"""
Interactive 3D globe view, rendered with vedo.

vedo and scipy are probed separately: scipy is a hard requirement of the rest of
the toolbox while vedo is optional, so a missing vedo must not be reported as a
missing scipy or vice versa.

The work is split three ways so that only `scene` needs a render window:
`geometry` is the sphere and its relief, `shell` is the atmosphere above it, and
`scene` is the window they are shown in.
"""

from dataclasses import dataclass

import numpy as np
import matplotlib.pyplot as plt
from matplotlib.colors import Normalize

from ...colors import TERRAIN_RANGE, truncated
from ...coords import close_longitude_seam, wrap_longitudes
from ...topography import load_topography
from . import geometry
from .geometry import (Globe, add_bar, apply_colormap, build_surface, densify,
                       surface_scalars)
from .scene import (CONTOUR_LIFT, GRATICULE_LIFT, LINE_STEP_DEG, GlobeScene,
                    subsolar_point)

# Where a terrain surface puts its elevation bar, leaving vedo's default
# right-hand slot to the shell above it.
SURFACE_BAR_POS = ((0.04, 0.25), (0.08, 0.75))

try:
    from vedo import Line, Mesh
    vedo_available = True
except ImportError:
    vedo_available = False

try:
    from scipy.interpolate import RegularGridInterpolator
    scipy_available = True
except ImportError:
    scipy_available = False


@dataclass
class GlobeOptions:
    """
    Presentation choices for the globe, kept together so the renderer signatures
    do not grow a parameter per feature.
    """
    output_path: str = None
    relief_mode: str = 'topo'     # 'topo' | 'none' | 'data' | a variable name
    relief: np.ndarray = None     # the relief field, on the same grid as the data
    relief_label: str = None
    sun: tuple = None             # explicit (lon, lat)
    sun_ls: float = None          # solar longitude, sets the sub-solar latitude
    export: str = None
    contours: bool = True
    lighting: bool = None         # None: lit only when a sub-solar point is given
    surface_mode: str = None      # 'field' | 'terrain' | 'none'; None picks one
    # Atmospheric shell
    shell_mode: str = 'cloud'     # 'cloud' | 'iso' | 'column' | 'layers'
    shell_threshold: str = None   # value(s), or 'pNN' for a percentile
    shell_color: str = 'value'    # 'value' | 'height'
    shell_level: int = None       # draw one level instead of the whole volume
    shell_opacity: float = 0.35
    shell_top_km: float = None    # where a non-metric vertical tops out
    shell_resolution: int = None  # voxels per side of the ray-cast box
    shell_cut: bool = None        # None: cut only the modes that need it
    shell_exaggeration: float = None  # vertical scale of the air; None: the default
    altitudes: object = None      # vertical.Altitudes, when the file gives them
    # Winds, as eastward/northward components on the same grid as the data
    wind_u: np.ndarray = None
    wind_v: np.ndarray = None
    # Animation
    frames: np.ndarray = None     # every step, the animated axis first
    frame_dim: str = None
    frame_axis: object = None     # the Axis the frames run over, for labelling
    fps: int = 12
    spin: float = None            # turns of camera rotation, None for none

    def lit(self, sun):
        """
        Whether to shade the globe.

        Flat by default: shading multiplies the colormap by the angle to the
        light, so a value reads differently depending on where it sits on the
        sphere, which is exactly what a colour scale is supposed to prevent.
        Asking for a sub-solar point is asking for the opposite - a lit planet
        with a terminator - so that turns it back on.
        """
        if self.lighting is not None:
            return bool(self.lighting)
        return sun is not None


def render_grid(nlat=180, nlon=360):
    """
    The 1 degree grid of cell centres the globe is meshed on.

    It matches the MOLA grid in topography.py, so a globe with relief and a
    globe without one are the same mesh with different radii.
    """
    lats = np.linspace(90 - 0.5, -90 + 0.5, nlat)
    lons = np.linspace(-180 + 0.5, 180 - 0.5, nlon)
    lon2d, lat2d = np.meshgrid(lons, lats)
    return lats, lons, lat2d, lon2d


def regrid(lat_src, lon_src, values, lat_grid, lon_grid):
    """
    Interpolate a lat/lon field onto the render grid.

    The incoming axes may be 2D, descending, or wrapped past 180 degrees, so
    both are normalized and sorted together with the data before use. Points the
    source does not cover come back NaN and are drawn as gaps.

    Duplicate columns are dropped. A map that had `add_cyclic_point` applied to
    an axis already spanning -180..180 carries a copy of the first column at
    191.25 degrees, which wraps straight back onto -168.75; the interpolator
    rejects a non-monotonic axis, so the copy has to go. Keeping the first
    occurrence is right because the copy holds the same data by construction.

    A field that goes all the way round then has its seam closed, the way the
    relief and the cloud volume close theirs. Without it the cells between the
    last longitude and the first are off the end of the axis and come back as
    gaps: a model grid of 32 longitudes ends at 168.75, so an eleven degree
    wedge of the globe was drawn as missing data - and missing data on a shell
    is opaque dark red, which is what the column mode's bad slice was.
    """
    lons = np.asarray(lon_src, dtype=float)
    lats = np.asarray(lat_src, dtype=float)
    values = np.asarray(values, dtype=float)

    lons = lons[0, :] if lons.ndim == 2 else lons
    lats = lats[:, 0] if lats.ndim == 2 else lats
    lons = wrap_longitudes(lons)

    # np.unique sorts and de-duplicates in one step, returning where each kept
    # value came from so the data can follow
    lats, lat_order = np.unique(lats, return_index=True)
    lons, lon_order = np.unique(lons, return_index=True)
    lons, values = close_longitude_seam(lons, values[lat_order, :][:, lon_order])

    interp = RegularGridInterpolator(
        (lats, lons),
        values,
        bounds_error=False,
        fill_value=np.nan,
    )
    return interp((lat_grid, lon_grid))


def topography_contours(globe, lat_grid, lon_grid, values, levels=10):
    """
    Topography contour segments, lifted onto the globe.

    Generated on a throwaway figure so the caller's current figure is left
    untouched.
    """
    tmp_fig = plt.figure()
    try:
        cs = tmp_fig.gca().contour(lon_grid, lat_grid, values, levels=levels, linewidths=0)
        segments = [seg for path in cs.get_paths()
                    for seg in path.to_polygons(closed_only=False)]
    finally:
        plt.close(tmp_fig)

    # Contours fall wherever the field puts them, so they cannot be snapped to
    # the grid the way the graticule is: a vertex sits on a cell edge and the
    # segment between two of them crosses the diagonal where the cell's two
    # triangles meet, cutting the corner off the ground. Splitting each segment
    # inside its cell is what makes the line lie in the terrain rather than over
    # it, and `on_surface` then reads the triangles themselves.
    return [Line(globe.on_surface(*densify(verts[:, 1], verts[:, 0], LINE_STEP_DEG),
                                  CONTOUR_LIFT))
            for verts in segments if len(verts) >= 2]


def graticule(globe, lats, lons):
    """
    Meridians every 30 degrees and parallels every 30 degrees, lying on the
    terrain.

    Sampled on the render grid's own coordinates rather than on a linspace of
    the same length. That sounds like a distinction without a difference and is
    not: the grid holds cell centres at half degrees, so a linspace over whole
    degrees puts every vertex in the middle of a cell, where the height has to
    be guessed and the guess can fall through the triangle it is guessing about.
    Landing on the nodes themselves means each vertex sits directly over a mesh
    vertex, and since the line and the mesh are both straight in between, a line
    that clears its own endpoints clears everything between them.
    """
    lines = []
    for lon in range(-150, 181, 30):
        meridian = np.full_like(lats, float(lon))
        lines.append(Line(globe.on_surface(lats, meridian, GRATICULE_LIFT)))
    # Stopping at 60: a "parallel" at the pole is a circle of no radius, drawn
    # as a ring of coincident points on top of the one place the grid cannot
    # answer for
    for lat in range(-60, 61, 30):
        parallel = np.full_like(lons, float(lat))
        # Closed back onto the first longitude, or a parallel stops one cell
        # short of the seam the surface itself wraps
        ring = np.append(lons, lons[0])
        lines.append(Line(globe.on_surface(np.append(parallel, parallel[0]),
                                           ring, GRATICULE_LIFT)))
    return lines


def _build_relief(options, lat_src, lon_src, data, lat_grid, lon_grid, norm=None):
    """
    (elevation in metres on the render grid, topography or None, note).

    'topo' is MOLA. 'none' is a bare sphere. Anything else shapes the globe by a
    field, which carries no metric height and so is normalized into a fixed
    fraction of the radius.
    """
    mode = options.relief_mode or 'topo'

    # Loaded whatever the relief is: the contours are a geographic reference,
    # useful even on a globe shaped by something other than the terrain.
    topo = load_topography()

    if mode == 'topo':
        if topo is None:
            return None, None, "topography missing, drawing a bare sphere"
        return topo.values, topo, None

    if mode == 'none':
        return np.zeros_like(lat_grid), topo, None

    # The plot's colour scale describes the plotted field, so it only applies to
    # the relief when the two are the same variable
    source = options.relief if options.relief is not None else data
    shared = norm if mode == 'data' else None
    field = regrid(lat_src, lon_src, source, lat_grid, lon_grid)
    label = options.relief_label or mode
    return (geometry.normalize_relief(field, norm=shared), topo,
            f"relief from '{label}', normalized to "
            f"{geometry.DATA_RELIEF_FRACTION:.0%} of the radius")


def plot_3D_globe(lon2d, lat2d, data2d, colormap, varname, units=None, norm=None,
                  options=None):
    """
    Plot a 3D globe view of the data using vedo, with surface coloring based on
    data2d and overlaid contour lines from topography.
    """
    if not vedo_available:
        print("3D view skipped: vedo missing.")
        return 1
    if not scipy_available:
        print("3D view skipped: scipy missing.")
        return 1

    options = options or GlobeOptions()
    lats, lons, lat_grid, lon_grid = render_grid()
    nlat, nlon = lat_grid.shape

    elevation, topo, note = _build_relief(options, lat2d, lon2d, data2d,
                                          lat_grid, lon_grid, norm=norm)
    if elevation is None:
        print(f"3D view skipped: {note}.")
        return 1
    if note:
        print(f"Note: {note}.")

    values = regrid(lat2d, lon2d, data2d, lat_grid, lon_grid)
    if not np.any(np.isfinite(values)):
        print("3D view skipped: the field does not cover any of the globe grid.")
        return 1

    sun = subsolar_point(options.sun, options.sun_ls)
    lit = options.lit(sun)

    globe = Globe(relief=elevation, lats=lats, lons=lons)
    mesh = _surface_mesh(globe, lat_grid, lon_grid, elevation, values,
                         colormap, norm, varname, units, lit=lit,
                         mode=options.surface_mode or 'field')

    scene = GlobeScene(globe, title="3D globe view")
    scene.surface = mesh
    scene.add(mesh)
    scene.add_lines(graticule(globe, lats, lons), color='k', width=1)
    if options.contours and topo is not None:
        scene.add_lines(topography_contours(globe, lat_grid, lon_grid, topo.values),
                        color='k', width=0.5)

    scene.add(_wind_arrows(globe, options, lat2d, lon2d))
    _attach_animation(scene, mesh, options, lat2d, lon2d, lat_grid, lon_grid,
                      colormap, norm, varname, units)
    # After the data animation, so the turn rides on top of it rather than
    # replacing it: the two answer different questions and compose
    if options.spin:
        scene.spin(options.spin, fps=options.fps)

    if sun is not None:
        scene.add_sun(*sun, lit=lit)

    scene.caption(_caption(varname, units, options, lit,
                           relief=note or f"relief: {options.relief_mode or 'topo'}"))
    scene.attach_readout(mesh, varname, units)

    if options.export:
        scene.export(options.export)

    return scene.show(output_path=options.output_path)


def _caption(varname, units, options, lit, relief=None, surface=None, shell=None):
    """
    The lines describing a globe view on screen.

    Every one of these is a choice the picture cannot show by itself: which
    variable, what shapes the sphere, what the shell is, whether it is lit.
    """
    lines = [varname + (f" [{units}]" if units else '')]
    if relief:
        lines.append(relief)
    if surface:
        lines.append(f"surface: {surface}")
    lines.extend(shell or [])
    lines.append(f"lighting {'on' if lit else 'off'}")
    return lines


def _attach_animation(scene, mesh, options, lat2d, lon2d, lat_grid, lon_grid,
                      colormap, norm, varname, units):
    """
    Re-colour the surface for each step of the animated dimension.

    The geometry never changes, so a frame is one regrid and one `cmap` call.
    Frames are regridded on demand rather than up front: the largest sample has
    669 steps, which pre-computed would be a third of a gigabyte.
    """
    frames = options.frames
    if frames is None or len(frames) < 2:
        return

    title = varname + (f' [{units}]' if units else '')

    # `choose_style` leaves the norm unset for an ordinary field, and an unset
    # norm makes vedo rescale to whatever it is handed - which for an animation
    # is one frame, so the colours would drift with the frame minimum and
    # maximum. Pin the limits over every frame instead.
    if norm is None:
        low, high = geometry.finite_range(frames)
        norm = Normalize(vmin=low, vmax=high)

    def on_frame(index):
        values = regrid(lat2d, lon2d, frames[index], lat_grid, lon_grid)
        apply_colormap(mesh, colormap, surface_scalars(values), norm, title)

    from ..movie import frame_label

    scene.animate(len(frames), on_frame,
                  label=options.frame_dim or 'frame', fps=options.fps,
                  describe=lambda i: frame_label(options.frame_axis,
                                                 options.frame_dim, i))


def _wind_arrows(globe, options, lat2d, lon2d, altitude=0.0):
    """
    Wind arrows for `--vector U,V --show-3d`, or None when no pair was given.
    """
    if options.wind_u is None or options.wind_v is None:
        return None
    from .shell import wind_arrows

    lats = lat2d[:, 0] if np.ndim(lat2d) == 2 else np.asarray(lat2d)
    lons = lon2d[0, :] if np.ndim(lon2d) == 2 else np.asarray(lon2d)
    return wind_arrows(globe, lats, lons, options.wind_u, options.wind_v,
                       altitude=altitude)


def render_globe(ctx):
    """
    Renderer for the `globe` plot kind.

    With a vertical axis this draws the atmospheric shell over the terrain; with
    only longitude and latitude it is the surface globe, the same view
    `--show-3d` offers after a map.
    """
    plan = ctx.plan
    options = ctx.globe or GlobeOptions()
    # The globe is the figure here rather than a view offered after one, so it
    # takes `output_path` itself instead of the '_globe' suffix extras.py uses
    options.output_path = ctx.output_path

    options.frames, options.frame_dim = ctx.frames, ctx.frame_dim
    options.frame_axis, options.fps = ctx.frame_axis, ctx.fps

    if plan.z is None:
        lon2d, lat2d = np.meshgrid(plan.x.values, plan.y.values)
        return plot_3D_globe(lon2d, lat2d, ctx.data, ctx.colormap, ctx.varname,
                             ctx.units, norm=ctx.norm, options=options)

    if not vedo_available:
        print("3D view skipped: vedo missing.")
        return 1

    from .shell import build_shell

    # The surface underneath is the field's own bottom level, so the globe shows
    # what the atmosphere is sitting on rather than an unrelated variable
    surface_values = ctx.data[0]
    lon2d, lat2d = np.meshgrid(plan.x.values, plan.y.values)

    lats, lons, lat_grid, lon_grid = render_grid()
    elevation, topo, note = _build_relief(options, lat2d, lon2d, surface_values,
                                          lat_grid, lon_grid, norm=ctx.norm)
    if elevation is None:
        elevation, topo = np.zeros_like(lat_grid), None
    elif note:
        print(f"Note: {note}.")

    sun = subsolar_point(options.sun, options.sun_ls)
    lit = options.lit(sun)
    globe = Globe(relief=elevation, lats=lats, lons=lons,
                  air_exaggeration=options.shell_exaggeration)

    def one_shell(varname, data, colormap, norm, units, bar_pos=None):
        return build_shell(
            globe, plan, data, colormap, varname, units, norm=norm,
            mode=options.shell_mode, threshold=options.shell_threshold,
            level=options.shell_level, opacity=options.shell_opacity,
            top_km=options.shell_top_km, color=options.shell_color, lit=lit,
            resolution=options.shell_resolution, cut=options.shell_cut,
            bar_pos=bar_pos, altitudes=options.altitudes)

    layers = _overlay_layers(ctx)
    count = len(layers) + 1
    shell = one_shell(ctx.varname, ctx.data, ctx.colormap, ctx.norm,
                      ctx.units, bar_pos=_shell_bar_pos(0, count))
    if shell.note:
        print(f"Shell: {shell.note}.")
    if not shell.actors:
        print("Note: the shell is empty; showing the surface only.")

    # Each --overlay variable becomes another atmosphere in its own colours.
    # Volumes composite correctly with one another, so two cloud decks read as
    # two decks where they are apart and as a mixture where they overlap.
    extra_shells = []
    for index, layer in enumerate(layers):
        other = one_shell(layer.varname, layer.data, layer.colormap, layer.norm,
                          layer.units, bar_pos=_shell_bar_pos(index + 1, count))
        if other.actors:
            extra_shells.append(other)
            print(f"Shell: {layer.varname} - {other.note}.")
        else:
            print(f"Warning: --overlay '{layer.varname}' has nothing to draw here.")

    # Under a shell the sphere goes neutral. Painting it with the same variable
    # and the same colormap as the atmosphere above it - which is what it used
    # to do, with only the bottom level - put two different things in the same
    # colours and left the whole view ambiguous.
    surface_mode = options.surface_mode or ('terrain' if shell.actors else 'field')
    surface = _surface_mesh(globe, lat_grid, lon_grid, elevation,
                            regrid(lat2d, lon2d, surface_values, lat_grid, lon_grid),
                            ctx.colormap, ctx.norm, ctx.varname, ctx.units,
                            lit=lit, mode=surface_mode,
                            under_shell=bool(shell.actors))

    scene = GlobeScene(globe, title="3D globe view")
    scene.surface = surface
    scene.shell = shell.actors[0] if shell.actors else None
    scene.add(surface, *shell.actors)
    for other in extra_shells:
        scene.add(*other.actors)
        scene.want_depth_peeling(other.depth_peeling)
    scene.want_depth_peeling(shell.depth_peeling)
    scene.add_lines(graticule(globe, lats, lons), color='k', width=1)
    if options.contours and topo is not None:
        scene.add_lines(topography_contours(globe, lat_grid, lon_grid, topo.values),
                        color='k', width=0.5)

    if sun is not None:
        scene.add_sun(*sun, lit=lit)

    # Last, so that the plane reaches the graticule and the contours too: a
    # cutaway with the lines still drawn across the opening is not a cutaway.
    # `build_shell` has already resolved --shell-cut against what the mode
    # itself needs, so `shell.cut` is the whole answer here.
    if shell.cut:
        scene.add_cutter(shell.slicer)
    if options.spin:
        scene.spin(options.spin, fps=options.fps)

    caption_lines = list(shell.caption)
    if extra_shells:
        from ...overlay import describe
        caption_lines.append(describe(ctx.composite))
    scene.caption(_caption(ctx.varname, ctx.units, options, lit,
                           relief=note, surface=surface_mode, shell=caption_lines))
    scene.attach_readout(surface, ctx.varname, ctx.units)

    if options.export:
        scene.export(options.export)
    return scene.show(output_path=options.output_path)


def _overlay_layers(ctx):
    """
    The --overlay layers of a composite, or nothing when there is no composite.
    """
    composite = getattr(ctx, 'composite', None)
    return composite.layers[1:] if composite else ()


# The right-hand column the shells' colour bars share, top to bottom.
BAR_COLUMN = (0.90, 0.94)
BAR_TOP = 0.78
BAR_BOTTOM = 0.06


def _shell_bar_pos(index, count):
    """
    Where the `index`-th of `count` shell colour bars goes.

    One bar takes the whole column and needs no arithmetic; several share it,
    because a composite has one scale per variable and they measure different
    things, so they cannot be merged into one bar.
    """
    if count < 2:
        return None
    span = (BAR_TOP - BAR_BOTTOM) / count
    top = BAR_TOP - span * index
    left, right = BAR_COLUMN
    # A tenth of each slot left as a gap, so two bars never touch
    return ((left, top - span * 0.9), (right, top))


# Greys for a globe that is only there to be the ground under something else.
# Neutral and stopping short of both ends, so nothing on it competes with the
# shell above it - see colors.TERRAIN_RANGE for why `bone` would not do.
TERRAIN_COLORMAP = truncated('gray', *TERRAIN_RANGE)


def _surface_mesh(globe, lat_grid, lon_grid, elevation, values,
                  colormap, norm, varname, units, lit=False, mode='field',
                  under_shell=False):
    """
    The solid globe, which a shell may then sit over.

    `mode` is what the sphere shows: `field` is the plotted variable, the
    ordinary case for a lat/lon map; `terrain` is the relief in greys, so the
    data colours belong to the shell alone; `none` is a flat grey ball.

    `under_shell` sends the bar to the left-hand slot whatever the mode, because
    the shell above takes the right-hand one. Deciding it from the mode alone,
    as this used to, left `--globe-surface field` under a shell with both bars
    in the same place.
    """
    if mode == 'terrain':
        scalars, colormap, norm = elevation / 1000.0, TERRAIN_COLORMAP, None
        title = 'elevation [km]'
    elif mode == 'none':
        scalars, title = None, ''
    else:
        scalars = values
        title = varname + (f' [{units}]' if units else '')

    pts, faces, mesh_scalars = build_surface(globe, lat_grid, lon_grid, elevation,
                                             scalars if scalars is not None
                                             else np.zeros_like(elevation))
    mesh = Mesh([pts, faces])

    if mode == 'none':
        mesh.c('grey6')
    else:
        bar_title = apply_colormap(mesh, colormap, mesh_scalars, norm, title)
        # Under a shell the surface's bar moves aside: the shell owns the
        # right-hand slot, because it is what the view is about.
        aside = under_shell or mode == 'terrain'
        add_bar(mesh, bar_title, SURFACE_BAR_POS if aside else None)

    mesh.compute_normals()
    mesh.lighting('default' if lit else 'off')
    return mesh
