"""
MOLA topography: loading, and the contour overlay drawn on latitude/longitude maps.

The array is loaded on first use rather than at import time, so importing this
package neither touches the filesystem nor prints anything.
"""

from dataclasses import dataclass
from functools import lru_cache

import numpy as np

from . import paths


@dataclass(frozen=True)
class Topography:
    """
    Topography on a regular 1 degree grid of cell centers: latitudes run
    89.5 -> -89.5 (descending) and longitudes -179.5 -> 179.5 (ascending).
    """
    values: np.ndarray   # (nlat, nlon)
    lats: np.ndarray     # (nlat,) descending
    lons: np.ndarray     # (nlon,) ascending
    lat2d: np.ndarray
    lon2d: np.ndarray

    @property
    def shape(self):
        return self.values.shape


@lru_cache(maxsize=1)
def load_topography():
    """
    Return the Topography, or None when the data file is missing.
    Cached, so the warning is printed at most once per run.
    """
    import os
    if not os.path.isfile(paths.TOPO_NPY):
        print(f"Warning: '{paths.TOPO_NPY}' not found! Topography contours disabled.")
        return None
    values = np.load(paths.TOPO_NPY)
    nlat, nlon = values.shape
    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 Topography(values=values, lats=lats, lons=lons, lat2d=lat2d, lon2d=lon2d)


def overlay_topography(ax, transform, levels=10, lon_0_360=False):
    """
    Overlay topography contours onto a given axes.
    Set lon_0_360 when the axes uses longitudes in [0, 360] rather than
    [-180, 180], which happens on plain (non-projected) axes.
    """
    topo = load_topography()
    if topo is None:
        return
    lon2d, lat2d, values = topo.lon2d, topo.lat2d, topo.values
    if lon_0_360:
        shifted = np.where(topo.lons < 0, topo.lons + 360.0, topo.lons)
        order = np.argsort(shifted)
        lon2d = np.tile(shifted[order], (values.shape[0], 1))
        values = values[:, order]
    ax.contour(
        lon2d, lat2d, values,
        levels=levels,
        linewidths=0.5,
        colors='black',
        transform=transform
    )
