#!/usr/bin/env python3
import glob
import os
import numpy as np
import pandas as pd

# ==========================
# USER SETTINGS
# ==========================
INPUT_DIR = "./et_compare_with_openet_polymean"
OUTPUT_CSV = "./et_compare_with_openet_polymean/et_metrics_vs_conus_summary.csv"

COL_GLOBAL = "global et"
COL_CONUS  = "conus et"
COL_OPENET = "openet et"

MIN_SAMPLES = 5

# Valid range: (0, 100]
ET_MIN_EXCLUSIVE = 0.0
ET_MAX_INCLUSIVE = 100.0

# ==========================
# Helpers
# ==========================
def norm_col(s: str) -> str:
    return "".join(str(s).strip().lower().split())

def find_col(df, wanted: str) -> str:
    wanted_n = norm_col(wanted)
    mapping = {norm_col(c): c for c in df.columns}
    if wanted_n not in mapping:
        raise ValueError(f"Missing required column '{wanted}'. Found columns: {list(df.columns)}")
    return mapping[wanted_n]

def filter_et(df, model_col, ref_col):
    """Keep only rows where both ETs are in (0, 100] and not NaN."""
    d = df[[model_col, ref_col]].dropna()
    mask = (
        (d[model_col] > ET_MIN_EXCLUSIVE) & (d[model_col] <= ET_MAX_INCLUSIVE) &
        (d[ref_col]   > ET_MIN_EXCLUSIVE) & (d[ref_col]   <= ET_MAX_INCLUSIVE)
    )
    return d.loc[mask]

def compute_metrics(model: np.ndarray, ref: np.ndarray):
    diff = model - ref
    bias = float(np.mean(diff))
    rmse = float(np.sqrt(np.mean(diff ** 2)))
    corr = float(np.corrcoef(model, ref)[0, 1]) if len(model) > 1 else np.nan
    return bias, rmse, corr

# ==========================
# Main
# ==========================
rows = []

for f in sorted(glob.glob(os.path.join(INPUT_DIR, "*.csv"))):
    site_id = os.path.splitext(os.path.basename(f))[0]
    df = pd.read_csv(f)

    c_global = find_col(df, COL_GLOBAL)
    c_conus  = find_col(df, COL_CONUS)
    c_openet = find_col(df, COL_OPENET)

    rec = {"site_id": site_id}

    # ---------- Global vs CONUS ----------
    v_g = filter_et(df, c_global, c_conus)
    n_g = len(v_g)
    if n_g >= MIN_SAMPLES:
        b, r, c = compute_metrics(v_g[c_global].to_numpy(), v_g[c_conus].to_numpy())
    else:
        b = r = c = np.nan

    rec.update({
        "n_global_vs_conus": n_g,
        "bias_global_vs_conus": b,
        "rmse_global_vs_conus": r,
        "corr_global_vs_conus": c
    })

    # ---------- OpenET vs CONUS ----------
    v_o = filter_et(df, c_openet, c_conus)
    n_o = len(v_o)
    if n_o >= MIN_SAMPLES:
        b, r, c = compute_metrics(v_o[c_openet].to_numpy(), v_o[c_conus].to_numpy())
    else:
        b = r = c = np.nan

    rec.update({
        "n_openet_vs_conus": n_o,
        "bias_openet_vs_conus": b,
        "rmse_openet_vs_conus": r,
        "corr_openet_vs_conus": c
    })

    rows.append(rec)

out = pd.DataFrame(rows).sort_values("site_id")
out.to_csv(OUTPUT_CSV, index=False)

print(f"? Wrote: {OUTPUT_CSV}")
print("?? ET filter applied: (0, 100] for BOTH model and CONUS benchmark")
