# -*- coding: utf-8 -*-
"""
Example: Read and visualize the Level-2 GIX NetCDF dataset.

This example prints metadata for every variable, reads every preset variable, extracts one time slice,
and plots example VTEC and GIX maps.
"""

import matplotlib.pyplot as plt
import netCDF4 as nc
import numpy as np


# ----------------------------------------------------------------------
# 1. Basic configuration
# ----------------------------------------------------------------------
filename = "GIX_BDSGEO_20242650000_01D_15M.nc"

# Select a time record by its zero-based position in the NetCDF ``time``
# coordinate:
#   0 = first stored time, 1 = second stored time, and so on.
# This is an array index, not an hour or a number of elapsed minutes.
#
# For a file beginning at 00:00 UTC with a 15-minute interval:
#   index 0  -> 00:00 UTC
#   index 1  -> 00:15 UTC
#   index 25 -> 06:15 UTC
target_time_index = 25

COORDINATE_VARIABLES = (
    "time", "lon_025deg", "lat_025deg", "lon_1deg", "lat_1deg"
)
DATA_VARIABLES = (
    "GIX", "GIX_std", "GIX_x", "GIX_y", "GIX_t_x", "GIX_t_y",
    "VTEC", "VTEC_t", "ROTI",
)
EXPECTED_VARIABLES = COORDINATE_VARIABLES + DATA_VARIABLES


def show_var_info(ds, var_name):
    """Print dimensions, data type, and all NetCDF attributes of a variable."""
    var = ds.variables[var_name]
    print(var_name)
    print(f"  Size       : {var.shape}")
    print(f"  Dimensions : {var.dimensions}")
    print(f"  Datatype   : {var.dtype}")
    print("  Attributes :")
    for attr in var.ncattrs():
        print(f"    {attr} = {getattr(var, attr)}")


def as_nan_array(var):
    """Read a NetCDF variable and convert masked/fill values to NaN."""
    return np.ma.asarray(var[:], dtype=np.float64).filled(np.nan)


def get_time_slice(var, t_index):
    """Extract a 2-D slice along a variable's named ``time`` dimension."""
    if "time" not in var.dimensions:
        raise ValueError(f"{var.name} does not contain a 'time' dimension")

    time_axis = var.dimensions.index("time")
    if not 0 <= t_index < var.shape[time_axis]:
        raise IndexError(
            f"target_time_index={t_index} is outside {var.name}'s "
            f"time-axis range 0..{var.shape[time_axis] - 1}"
        )

    return np.take(as_nan_array(var), indices=t_index, axis=time_axis)


# ----------------------------------------------------------------------
# 2. Inspect and read every variable preset by main_GIX_ver2.0.py
# ----------------------------------------------------------------------
with nc.Dataset(filename, mode="r") as ds:
    print("====== Dataset summary ======")
    print(ds)

    print("\n====== Variable list ======")
    print(list(ds.variables.keys()))

    missing_variables = [name for name in EXPECTED_VARIABLES if name not in ds.variables]
    if missing_variables:
        raise KeyError(f"Missing expected variables: {missing_variables}")

    print("\n====== Variable info ======")
    for variable_name in ds.variables:
        show_var_info(ds, variable_name)

    data = {name: as_nan_array(ds.variables[name]) for name in EXPECTED_VARIABLES}
    slices = {
        name: get_time_slice(ds.variables[name], target_time_index)
        for name in DATA_VARIABLES
    }


# Coordinate variables
time = data["time"]                  # seconds since midnight UTC
lon_025deg = data["lon_025deg"]     # GIX-family longitude grid
lat_025deg = data["lat_025deg"]     # GIX-family latitude grid
lon_1deg = data["lon_1deg"]         # VTEC/VTEC_t/ROTI longitude grid
lat_1deg = data["lat_1deg"]         # VTEC/VTEC_t/ROTI latitude grid

# All full 3-D variables remain available in ``data``. All selected 2-D
# fields, including GIX_std, are available in ``slices``.
GIX_slice = slices["GIX"]
VTEC_slice = slices["VTEC"]

selected_time_seconds = time[target_time_index]
selected_time_hours = selected_time_seconds / 3600.0


# ----------------------------------------------------------------------
# 3. Example visualization: VTEC (1 degree) and GIX (0.25 degree)
# ----------------------------------------------------------------------
plt.rcParams.update({"font.size": 14})
fig, axes = plt.subplots(1, 2, figsize=(12, 5))

ax = axes[0]
im_vtec = ax.pcolormesh(
    lon_1deg, lat_1deg, VTEC_slice, shading="auto", cmap="jet"
)
im_vtec.set_clim(0, 150)
cbar_vtec = fig.colorbar(im_vtec, ax=ax)
cbar_vtec.set_label("VTEC (TECU)", fontsize=16)
cbar_vtec.ax.tick_params(labelsize=14)

ax.plot(
    [95, 135, 135, 95, 95], [15, 15, 50, 50, 15],
    color=[0.5, 0.5, 0.5], linewidth=1
)
ax.set_xlim(lon_1deg[0], lon_1deg[-1])
ax.set_ylim(lat_1deg[0], lat_1deg[-1])
ax.set_xticks(np.arange(lon_1deg[0], lon_1deg[-1] + 1, 10))
ax.set_yticks(np.arange(lat_1deg[0], lat_1deg[-1] + 1, 5))
ax.grid(True, linestyle="--", alpha=0.5)
ax.set_xlabel("Longitude (deg)", fontsize=16)
ax.set_ylabel("Latitude (deg)", fontsize=16)
ax.set_title(f"VTEC at {selected_time_hours:g} UT", fontsize=18)
ax.tick_params(
    axis="both", which="major", direction="out",
    length=6, width=1, labelsize=14
)

ax = axes[1]
im_gix = ax.pcolormesh(
    lon_025deg, lat_025deg, GIX_slice, shading="auto", cmap="jet"
)
im_gix.set_clim(0, 100)
cbar_gix = fig.colorbar(im_gix, ax=ax)
cbar_gix.set_label(r"GIX ($10^{-3}$ TECU/km)", fontsize=16)
cbar_gix.ax.tick_params(labelsize=14)

ax.plot(
    [95, 135, 135, 95, 95], [15, 15, 50, 50, 15],
    color=[0.5, 0.5, 0.5], linewidth=1
)
ax.set_xlim(lon_025deg[0], lon_025deg[-1])
ax.set_ylim(lat_025deg[0], lat_025deg[-1])
ax.set_xticks(np.arange(lon_025deg[0], lon_025deg[-1] + 1, 10))
ax.set_yticks(np.arange(lat_025deg[0], lat_025deg[-1] + 1, 5))
ax.grid(True, linestyle="--", alpha=0.5)
ax.set_xlabel("Longitude (deg)", fontsize=16)
ax.set_ylabel("Latitude (deg)", fontsize=16)
ax.set_title(f"GIX at {selected_time_hours:g} UT", fontsize=18)
ax.tick_params(
    axis="both", which="major", direction="out",
    length=6, width=1, labelsize=14
)

fig.tight_layout()
savename = "VTEC_GIX_example.png"
fig.savefig(savename, dpi=300, bbox_inches="tight")
plt.show()

print(f"Figure saved as: {savename}")
