#!/usr/bin/env python3

import xarray as xr
import matplotlib.pyplot as plt

# --------------------------------------------------
# Input file
# --------------------------------------------------

FILE = "../202402_ONERA_NA_domain.nc4"

# --------------------------------------------------
# Read dataset
# --------------------------------------------------

ds = xr.open_dataset(FILE)

# --------------------------------------------------
# Create figure
# --------------------------------------------------

fig, axes = plt.subplots(
    1, 3,
    figsize=(18, 6),
    constrained_layout=True
)

# --------------------------------------------------
# 1) Horizontal flight probability distribution
# --------------------------------------------------

planes = ds["Number_of_kms"].sel(level=225, method="nearest").sum(dim="time")
prob = planes / planes.sum()

print(f"Horizontal probability sum = {float(prob.sum()):.6f}")

prob.plot(
    ax=axes[0],
    x="lon",
    y="lat",
    cmap="viridis",
    robust=True,
    cbar_kwargs={
        "label": "Flight probability"
    }
)

axes[0].set_title(
    "Horizontal flight probability"
)
axes[0].set_xlabel("Longitude")
axes[0].set_ylabel("Latitude")


# --------------------------------------------------
# 2) Vertical flight probability distribution
# --------------------------------------------------

planes = ds["Number_of_kms"].sum(dim=("time", "lat", "lon"))
prob = planes / planes.sum()

print(f"Vertical probability sum = {float(prob.sum()):.6f}")

prob.plot(
    ax=axes[1],
    y="level",
    marker="o"
)

axes[1].invert_yaxis()
axes[1].set_xlabel("Probability")
axes[1].set_ylabel("Pressure (hPa)")
axes[1].set_title(
    "Vertical flight probability"
)
axes[1].grid(True)


# --------------------------------------------------
# 3) Latitude-pressure distribution
# --------------------------------------------------

planes = ds["Number_of_kms"].sum(dim=("time", "lon"))
prob = planes / planes.sum()

print(f"Latitude-pressure probability sum = {float(prob.sum()):.6f}")

prob.plot(
    ax=axes[2],
    x="lat",
    y="level",
    cmap="viridis",
    robust=True,
    cbar_kwargs={
        "label": "Flight probability"
    }
)

axes[2].invert_yaxis()
axes[2].set_xlabel("Latitude")
axes[2].set_ylabel("Pressure (hPa)")
axes[2].set_title(
    "Latitude–pressure flight probability"
)


# --------------------------------------------------
# Save figure
# --------------------------------------------------

plt.savefig(
    "flight_probability_distribution.png",
    dpi=200,
    bbox_inches="tight"
)

plt.show()
