#!/usr/bin/env python3
"""
NMME Forecast Figure Generator
================================
Walks a data directory tree, reads NetCDF files, and generates
map figures. Also writes a log of generated files to manifest.log.
"""

import os
import re
import json
import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from pathlib import Path
import datetime

# ── USER CONFIG ───────────────────────────────────────────────────────────────
DATA_DIR   = "/cpc/home/jinfanti/MET/S2S/nmme_dev/nmme_met_development/output1"
OUTPUT_DIR = "/cpc/home/jinfanti/MET/S2S/nmme_dev/nmme_met_development/final/plots_emily/figures"
DPI        = 120

# Explicit Filters
TARGET_EXPERIMENTS = ["Beta1.1", "CFSv2"]
TARGET_YEARS       = "1991_2020"
TARGET_IC_MONTHS = [datetime.datetime.now().strftime("%m")]

# DIFFERENCE PLOTTING CONFIG
COMPUTE_DIFF = True
DIFF_PAIR    = ("Beta1.1", "CFSv2")  # Will compute (Beta1.0 - CFSv2)
# ─────────────────────────────────────────────────────────────────────────────

try:
    import netCDF4 as nc
except ImportError:
    raise ImportError("Please install netCDF4:  pip install netCDF4")

try:
    import cartopy.crs as ccrs
    import cartopy.feature as cfeature
    HAS_CARTOPY = True
except ImportError:
    print("WARNING: cartopy not found – falling back to plain matplotlib maps.")
    HAS_CARTOPY = False

# ── MONTH MAP ─────────────────────────────────────────────────────────────────
IC_MONTH_NAMES = {
    "01":"January","02":"February","03":"March","04":"April",
    "05":"May","06":"June","07":"July","08":"August",
    "09":"September","10":"October","11":"November","12":"December",
}

# ── WHICH netCDF variables to plot for each metric ────────────────────────────
METRIC_VARS = {
    "AnomCorr": ["series_cnt_PR_CORR"],
    "Bias":     ["series_cnt_ME", "series_cnt_RMSE", "series_cnt_MAE"],
    "BSS":      ["series_pstd_BSS"],
    "HSS":      ["series_cnt_HSS"],
}

# ── COLORMAP / RANGE per metric + variable ────────────────────────────────────
VAR_STYLE = {
    "series_cnt_PR_CORR": dict(cmap="RdBu_r",  vmin=-1.0, vmax=1.0, label="Pearson Correlation", extend="both"),
    "series_cnt_ME":      dict(cmap="BrBG",    vmin=-2.0, vmax=2.0, label="Mean Error", extend="both"),
    "series_cnt_RMSE":    dict(cmap="YlOrRd",  vmin=0.0,  vmax=3.0, label="RMSE", extend="max"),
    "series_cnt_MAE":     dict(cmap="YlOrRd",  vmin=0.0,  vmax=3.0, label="MAE", extend="max"),
    "series_pstd_BSS":    dict(cmap="RdYlGn",  vmin=-0.5, vmax=0.5, label="Brier Skill Score", extend="both"),
    "series_cnt_HSS":     dict(cmap="RdYlGn",  vmin=-0.5, vmax=1.0, label="Heidke Skill Score", extend="both"),
}
DEFAULT_STYLE = dict(cmap="RdBu_r", vmin=-1.0, vmax=1.0, label="Value", extend="both")

VAR_NAMES = {
    "prate":  "Precipitation Rate",
    "tmp2m":  "2-m Temperature",
    "tmpsfc": "Surface Temperature",
}

# ── FILENAME PATTERN ──────────────────────────────────────────────────────────
FNAME_RE = re.compile(
    r"(?P<metric>AnomCorr|Bias|BSS|HSS)"   
    r"_SeriesAnalysis_"
    r"(?P<variable>prate|tmp2m|tmpsfc)"     
    r"_(?P<vartype>[^_]+(?:_[^_]+)*?)_"    
    r"(?P<experiment>.+?)"                  
    r"_IC(?P<ic>\d{2})"                     
    r"_Lead(?P<lead>\d+)seasonal"           
    r"_(?P<yr1>\d{4})_(?P<yr2>\d{4})"      
    r"(?:_(?P<condition>[^.]+))?"           
    r"\.nc$",
    re.IGNORECASE
)

EXP_NORM = {
    "nmmewithsfs":      "NMMEwithSFS",
    "nmmewithsfsbeta":  "NMMEwithSFSBeta",
    "nmmewithcfs":      "NMMEwithCFS",
    "cfsv2":            "CFSv2",
    "sfs_baseline":     "SFS_Baseline",
    "sfs_v0_1_100":     "SFS_v0.1_100",
    "beta1.0":          "Beta1.0",
    "beta1.1":          "Beta1.1",
}

def normalise_experiment(raw):
    return EXP_NORM.get(raw.lower(), raw)

# ── PARSE FILENAME ────────────────────────────────────────────────────────────
def parse_filename(fname):
    m = FNAME_RE.match(os.path.basename(fname))
    if not m:
        return None
    d = m.groupdict()
    d["lead"]       = str(int(d["lead"]))          
    d["ic"]         = d["ic"].zfill(2)             
    d["ic_name"]    = IC_MONTH_NAMES.get(d["ic"], d["ic"])
    d["var_name"]   = VAR_NAMES.get(d["variable"], d["variable"])
    d["experiment"] = normalise_experiment(d["experiment"])
    d["metric"]     = d["metric"].capitalize()     
    
    d["condition"]  = d["condition"] if d["condition"] else "All"

    for k in METRIC_VARS:
        if k.lower() == d["metric"].lower():
            d["metric"] = k
            break
    return d


# ── LOAD DATA ─────────────────────────────────────────────────────────────────
def load_data(filepath, meta):
    ds = nc.Dataset(filepath, "r")
    try:
        lats = ds.variables["lat"][:]
        lons = ds.variables["lon"][:]

        wanted   = METRIC_VARS.get(meta["metric"], [])
        skip     = {"lat", "lon", "n_series"}
        all_2d   = [v for v in ds.variables
                    if v not in skip and ds.variables[v].ndim == 2]

        results = []
        for nc_var in all_2d:
            clean = re.sub(r'\d+$', '', nc_var)
            match_key = None
            for w in wanted:
                if nc_var == w or clean == w:
                    match_key = w
                    break
            if match_key is None and not wanted:
                match_key = nc_var   

            if match_key is None:
                continue

            raw  = ds.variables[nc_var][:]
            fill = getattr(ds.variables[nc_var], "_FillValue", -9999.0)
            data = np.ma.masked_equal(raw, fill)
            data = np.ma.masked_invalid(data)
            results.append((match_key, lons, lats, data))

    finally:
        ds.close()

    return results


# ── PLOTTING ──────────────────────────────────────────────────────────────────
def build_title(meta, nc_var, is_diff=False):
    style  = VAR_STYLE.get(nc_var, DEFAULT_STYLE)
    cond   = (meta["condition"]
              .replace("ElNino","El Niño")
              .replace("LaNina","La Niña")
              .replace("LowerTercile","Lower Tercile")
              .replace("UpperTercile","Upper Tercile"))
    vtype  = meta["vartype"].replace("_"," ")
    
    label = f"Difference ({style['label']})" if is_diff else style['label']
    
    return (
        f"{meta['experiment']}  |  {meta['var_name']} ({vtype})  |  "
        f"{meta['metric']} – {label}\n"
        f"IC: {meta['ic_name']}   Lead: {meta['lead']}   "
        f"Condition: {cond}   ({meta['yr1']}–{meta['yr2']})"
    )


def make_figure(lons, lats, data, meta, nc_var, out_path, is_diff=False):
    style = dict(VAR_STYLE.get(nc_var, DEFAULT_STYLE)) 
    
    # ── DIFFERENCE PLOT STYLING ──────────────────────────────────────────────
    if is_diff:
        # 1. For Errors (RMSE, MAE): Lower is better. 
        # Diff < 0 means Beta1.0 is better. 
        # RdBu_r makes negative values Blue (Beta better) and positive Red (CFS better).
        if any(m in nc_var for m in ["RMSE", "MAE"]):
            style["cmap"] = "RdBu_r"
            style["vmin"], style["vmax"] = -1.0, 1.0
            style["extend"] = "both"
            
        # 2. For Skill Scores (PR_CORR, BSS, HSS): Higher is better. 
        # Diff > 0 means Beta1.0 is better.
        # RdBu makes positive values Blue (Beta better) and negative Red (CFS better).
        elif any(m in nc_var for m in ["PR_CORR", "BSS", "HSS"]):
            style["cmap"] = "RdBu"
            style["vmin"], style["vmax"] = -0.5, 0.5  
            style["extend"] = "both"
            
        # 3. For Mean Error (Bias): 
        # Diff just shows relative shift, not absolute skill. Use neutral diverging.
        else: 
            style["cmap"] = "PuOr" 
            style["vmin"], style["vmax"] = -1.0, 1.0
            style["extend"] = "both"
    # ─────────────────────────────────────────────────────────────────────────

    if HAS_CARTOPY:
        fig = plt.figure(figsize=(12, 6))
        ax  = fig.add_subplot(1, 1, 1, projection=ccrs.PlateCarree(central_longitude=180))
        ax.set_global()
        lons_plot = lons - 180.0
        im = ax.pcolormesh(lons_plot, lats, data,
                           transform=ccrs.PlateCarree(central_longitude=180),
                           cmap=style["cmap"], vmin=style["vmin"], vmax=style["vmax"],
                           shading="auto")
        ax.add_feature(cfeature.COASTLINE, linewidth=0.6, edgecolor="k")
        ax.add_feature(cfeature.BORDERS,   linewidth=0.3, edgecolor="0.4")
        ax.gridlines(draw_labels=False, linewidth=0.3, color="gray", alpha=0.5)
    else:
        fig, ax = plt.subplots(figsize=(12, 6))
        im = ax.pcolormesh(lons, lats, data,
                           cmap=style["cmap"], vmin=style["vmin"], vmax=style["vmax"],
                           shading="auto")
        ax.set_xlabel("Longitude"); ax.set_ylabel("Latitude")

    cb = plt.colorbar(im, ax=ax, orientation="horizontal",
                      pad=0.04, fraction=0.046, extend=style["extend"])
    
    if is_diff and "ME" not in nc_var:
        label_text = f"Diff ({style['label']})  [Blue = Beta1.0 Better]"
    else:
        label_text = style["label"] if not is_diff else f"Diff ({style['label']})"
        
    cb.set_label(label_text, fontsize=10, fontweight="bold")
    ax.set_title(build_title(meta, nc_var, is_diff), fontsize=10, fontweight="bold", pad=8)

    plt.tight_layout()
    plt.savefig(out_path, dpi=DPI, bbox_inches="tight")
    plt.close(fig)


# ── OUTPUT PATH ───────────────────────────────────────────────────────────────
def get_output_path(base_dir, meta, nc_var):
    Path(base_dir).mkdir(parents=True, exist_ok=True)
    stat  = nc_var.split("_")[-1]  
    fname = (f"{meta['experiment']}_{meta['variable']}_{meta['metric']}_{stat}"
             f"_IC{meta['ic']}_Lead{meta['lead'].zfill(2)}"
             f"_{meta['vartype']}_{meta['condition']}.png").lower()
    return str(Path(base_dir) / fname)


# ── MAIN ─────────────────────────────────────────────────────────────────────
def main():
    data_root = Path(DATA_DIR)
    if not data_root.exists():
        raise FileNotFoundError(f"DATA_DIR not found: {DATA_DIR}")
    Path(OUTPUT_DIR).mkdir(parents=True, exist_ok=True)

    nc_files = list(data_root.rglob("*.nc"))
    print(f"Found {len(nc_files)} NetCDF files under {DATA_DIR}\n")

    manifest = []
    saved = skipped = errors = diffs_computed = 0
    
    # 1. First Pass: Group files by identical parameters (except experiment)
    groups = {}
    for fpath in nc_files:
        meta = parse_filename(str(fpath))
        if not meta:
            continue
            
        if meta["experiment"] not in TARGET_EXPERIMENTS:
            continue
        if f"{meta['yr1']}_{meta['yr2']}" != TARGET_YEARS:
            continue
        if meta["ic"] not in TARGET_IC_MONTHS:
            continue
        if meta["condition"] in ("ElNino(1)", "LaNina(1)"):
            continue

        # Create a unique signature for the file context
        signature = (meta["variable"], meta["metric"], meta["vartype"], 
                     meta["ic"], meta["lead"], meta["condition"], meta["yr1"], meta["yr2"])
                     
        if signature not in groups:
            groups[signature] = {}
        groups[signature][meta["experiment"]] = (fpath, meta)

    # 2. Second Pass: Process groups (Plot individuals and compute differences)
    for signature, experiments in groups.items():
        datasets_cache = {}
        
        # Plot individual files first
        for exp_name, (fpath, meta) in experiments.items():
            try:
                datasets = load_data(str(fpath), meta)
                datasets_cache[exp_name] = datasets
            except Exception as e:
                print(f"ERROR reading {fpath.name}: {e}")
                errors += 1
                continue

            for nc_var, lons, lats, data in datasets:
                out = get_output_path(OUTPUT_DIR, meta, nc_var)
                stat = nc_var.split("_")[-1]

                if not Path(out).exists():
                    print(f"Plotting [{stat}] {fpath.name} …")
                    make_figure(lons, lats, data, meta, nc_var, out, is_diff=False)
                    saved += 1

                manifest.append({
                    "experiment": meta["experiment"],
                    "variable":   meta["variable"],
                    "var_name":   meta["var_name"],
                    "metric":     meta["metric"],
                    "stat":       stat,
                    "vartype":    meta["vartype"],
                    "ic":         meta["ic"],
                    "ic_name":    meta["ic_name"],
                    "lead":       meta["lead"],
                    "condition":  meta["condition"],
                    "yr1":        meta["yr1"],
                    "yr2":        meta["yr2"],
                    "path":       Path(out).name,
                })

        # Plot differences if both experiments are present in this group
        if COMPUTE_DIFF and DIFF_PAIR[0] in datasets_cache and DIFF_PAIR[1] in datasets_cache:
            exp1_data = datasets_cache[DIFF_PAIR[0]]
            exp2_data = datasets_cache[DIFF_PAIR[1]]
            meta_base = experiments[DIFF_PAIR[0]][1].copy() # Use exp1 meta as template
            
            diff_exp_name = f"{DIFF_PAIR[0]}-{DIFF_PAIR[1]}"
            meta_base["experiment"] = diff_exp_name

            # Match variables between the two files
            for nc_var1, lons1, lats1, data1 in exp1_data:
                for nc_var2, lons2, lats2, data2 in exp2_data:
                    if nc_var1 == nc_var2:
                        # Compute difference
                        diff_data = data1 - data2
                        out = get_output_path(OUTPUT_DIR, meta_base, nc_var1)
                        stat = nc_var1.split("_")[-1]
                        
                        if not Path(out).exists():
                            print(f"Plotting [{stat}] DIFFERENCE {diff_exp_name} …")
                            make_figure(lons1, lats1, diff_data, meta_base, nc_var1, out, is_diff=True)
                            saved += 1
                            diffs_computed += 1

                        manifest.append({
                            "experiment": meta_base["experiment"],
                            "variable":   meta_base["variable"],
                            "var_name":   meta_base["var_name"],
                            "metric":     meta_base["metric"],
                            "stat":       stat,
                            "vartype":    meta_base["vartype"],
                            "ic":         meta_base["ic"],
                            "ic_name":    meta_base["ic_name"],
                            "lead":       meta_base["lead"],
                            "condition":  meta_base["condition"],
                            "yr1":        meta_base["yr1"],
                            "yr2":        meta_base["yr2"],
                            "path":       Path(out).name,
                        })

    # Write manifest to a log file instead of JSON/HTML injection
    log_path = Path(OUTPUT_DIR) / "manifest.log"
    with open(log_path, "w") as f:
        json.dump(manifest, f, indent=2)

    print(f"\n✓ Done.  {saved} new figures plotted ({diffs_computed} were differences) |  {errors} errors")
    print(f"✓ Manifest logged to → {log_path}")

if __name__ == "__main__":
    main()
