"""
Functions to generate maps, and cross sections through tomographic models

Copyright 2023 Lars Gebraad and Thomas Schouten
"""

# Import packages
import numpy as _numpy
import pandas as _pandas
import xarray as _xarray
import cartopy.crs as ccrs
import pygmt as _pygmt

#--------------------#
#        MAPS        #
#--------------------#

def depth_slice(ax, fig, model, depth, plotting_options):
    """
    Function to plot depth slice from tomography model
    """
    # Create basemap
    ax, gl = basemap(ax, plotting_options["coastlines"])

    # Select depth slice and remove nan values
    data = model.sel(depth=depth, method="nearest")
    
    # Plot depth slice
    im = ax.imshow(
        data,
        transform=ccrs.PlateCarree(),
        cmap=plotting_options["colourmap"],
        vmax=plotting_options["max_value"],
        vmin=plotting_options["min_value"],
        extent=plotting_options["extent"]
    )

    # Add colourbar
    if plotting_options["colourbar"]:
        cbar = fig.colorbar(im, ax=ax, label=plotting_options["colourbar_label"], orientation=plotting_options["colourbar_orientation"], shrink=0.8, aspect=20)

    return ax

def basemap(ax, coastlines):
    """
    Function to plot basemap to display tomographic models
    """

    # Plot coastlines
    if coastlines:
        ax.coastlines(lw=0.5)

    # Set labels
    # ax.set_xlabel("Longitude")
    # ax.set_ylabel("Latitude")

    # Set global extent
    ax.set_global()

    # Set gridlines
    gl = ax.gridlines(
        crs=ccrs.PlateCarree(), 
        draw_labels=True, 
        linewidth=0.5, 
        color='gray', 
        alpha=0.5, 
        linestyle='--', 
        zorder=10
    )

    gl.top_labels = False
    gl.right_labels = False  

    return ax, gl

def slice_on_globe(ax, fig, point_A: _numpy.array, point_B: _numpy.array, plotting_options, whole_globe=False):
    """
    Function to plot slice through REVEAL on globe
    """

    # Check if point A and point B are properly defined
    assert len(point_A) == 2
    assert len(point_B) == 2

    points_on_great_circle = linspace_between_coordinates(
        point_A, point_B, plotting_options["number_of_points"], whole_globe=whole_globe
    )

    ax.plot(
        *_numpy.stack([point_A, point_B]).T,
        transform=ccrs.Geodetic(),
        linestyle="-",
        c="k",
        lw=1
    )

    ax.scatter(
        points_on_great_circle.r,
        points_on_great_circle.s,
        transform=ccrs.Geodetic(),
        zorder=1000,
        s=30,
        c=["k"] + (plotting_options["number_of_points"] - 2) * [plotting_options["fill_color"]] + ["w"],
        edgecolors="k",
    )

    # Generate dummy invisible colorbar for plt.tight_layout()
    if plotting_options["colourmap"] is True:
        data = _numpy.random.random((10, 10))
        p = ax.imshow(data)

        cbar = fig.colorbar(p, ax=ax)
        cbar.outline.set_visible(False)
        cbar.ax.set_visible(False)

    return ax

#--------------------#
#       ANNULI       #
#--------------------#

def slice(
    ax,
    fig,
    data: _xarray.DataArray,
    point_A: _numpy.array,
    point_B: _numpy.array,
    plotting_options
):
    
    # Check if data is already a slice or the whole model
    if data.dims == ("depth", "azimuthal_distance"):
        data_slice = data
    else:
        data_slice = extract_slice(data, point_A, point_B, plotting_options["azimuthal_pixels"])

    azimuthal_distance_degrees_grid, depth_grid = _numpy.meshgrid(
        data_slice.azimuthal_distance, data_slice.depth
    )
    
    azimuthal_distance_radians_grid = _numpy.deg2rad(azimuthal_distance_degrees_grid)
    radius_grid = 6371 - depth_grid

    if radius_grid.shape == data_slice.T.shape:
        data_slice = data_slice.T

    # Plot wavespeed anomaly
    contours = ax.contourf(
        azimuthal_distance_radians_grid,
        radius_grid,
        data_slice,
        cmap=plotting_options["colourmap"],
        levels=_numpy.arange(plotting_options["min_value"], plotting_options["max_value"], plotting_options["contour_interval"]),
        extend="both",
    )

    angle_between_points = calculate_great_circle_angle(point_A, point_B)

    ax.set_rorigin(1)  # Op hoop van zegen
    ax.set_theta_offset(-_numpy.deg2rad(180 + 90 - angle_between_points / 2))

    assert plotting_options["number_of_points"] >= 3
    indicator_points = linspace_between_coordinates(
        point_A, point_B, plotting_options["number_of_points"]
    )

    # Plot points along cross section
    ax.scatter(
        _numpy.deg2rad(indicator_points.p),
        _numpy.ones_like(indicator_points.p) * radius_grid.max(),
        zorder=1000,
        s=60,
        c=["k"] + (plotting_options["number_of_points"] - 2) * [plotting_options["fill_color"]] + ["w"],
        edgecolors="k",
        clip_on=False,
    )

    azimuths = _numpy.deg2rad(_numpy.sort(data_slice.azimuthal_distance))
    for interface in plotting_options["transition_zones"]:
        ax.plot(
            azimuths, 6371 - interface * _numpy.ones_like(azimuths), "--k", lw=0.5
        )

    ax.set_theta_direction(-1)

    ax.set_xlim(
        [
            _numpy.deg2rad(angle_between_points),
            _numpy.deg2rad(0),
        ]
    )
    ax.set_xticks([])
    ax.set_yticks([])

    if plotting_options["colourbar"]:
        cbar = fig.colorbar(contours, ax=ax, orientation="horizontal", shrink=0.6, aspect=20)
        cbar.ax.set_xlabel("dV/V [%]")

    return ax

def extract_slice(
    data: _xarray.DataArray,
    point_A: _numpy.array,
    point_B: _numpy.array,
    azimuthal_pixels: int = 1000,
):
    """
    Function to extract slice from REVEAL
    """
    points = linspace_between_coordinates(point_A, point_B, azimuthal_pixels)
    
    latitude = _xarray.DataArray(points.s, dims="azimuthal_distance")
    longitude = _xarray.DataArray(points.r, dims="azimuthal_distance")

    data_slice = data.load().interp(
        latitude=latitude, longitude=longitude, method="linear"
    )
    data_slice = data_slice.assign_coords(
        {
            "azimuthal_distance": calculate_great_circle_angle(point_A, point_B)
            * _numpy.linspace(0, 1, azimuthal_pixels)
        }
    )
    return data_slice

def linspace_between_coordinates(point_A, point_B, total_points, whole_globe=False):
    if whole_globe:
        azimuth = calculate_azimuth(point_A, point_B)

        return _pygmt.project(
            center=point_A, azimuth=azimuth, generate=360. / (total_points - 1), length="0/360", width="0/1"
        )
    
    else:
        angle = calculate_great_circle_angle(point_A, point_B)

        return _pygmt.project(
            center=point_A, endpoint=point_B, generate=angle / (total_points - 1)
        )

def calculate_great_circle_angle(point_A, point_B):
    """
    Function to calculate great circle angle between two points
    """
    coordinates = _numpy.deg2rad(_numpy.stack([point_A, point_B]))

    longitudes = coordinates[:, 0]
    latitudes = coordinates[:, 1]

    return _numpy.rad2deg(
        _numpy.arccos(
            _numpy.prod(_numpy.sin(latitudes))
            + _numpy.prod(_numpy.cos(latitudes)) * _numpy.cos(_numpy.diff(longitudes))
        )
    ).item()

def calculate_azimuth(point_A, point_B):
    """
    Function to calculate the azimuth between two points.
    """
    lat_A, lon_A = _numpy.deg2rad(point_A[0]), _numpy.deg2rad(point_A[1])
    lat_B, lon_B = _numpy.deg2rad(point_B[0]), _numpy.deg2rad(point_B[1])

    dlon = lon_B - lon_A

    x = _numpy.sin(dlon) * _numpy.cos(lat_B)
    y = _numpy.cos(lat_A) * _numpy.sin(lat_B) - (_numpy.sin(lat_A) * _numpy.cos(lat_B) * _numpy.cos(dlon))

    return _numpy.rad2deg(
        (_numpy.arctan2(x, y) + 2 * _numpy.pi) % (2 * _numpy.pi)
    )
