#!/usr/bin/env python3

import os
import glob
import numpy as np
import xarray as xr
import dask
from tqdm import tqdm
from model import Model
import inout as io
import gridop as gop

os.environ["HDF5_USE_FILE_LOCKING"] = "FALSE"

def main():
    input_dir = "/path/to/CROCO_FILES"
    output_dir = os.path.join(input_dir, "Z_CROCO")

    os.makedirs(output_dir, exist_ok=True)

    grid_master = os.path.join(input_dir, "croco_avg_72.00000.nc")
    file_list = sorted(glob.glob(os.path.join(input_dir, "croco_avg_*.nc")))
    print(f"Total files discovered: {len(file_list)}")

    tgrid = np.concatenate([
        np.linspace(0, -100, 10),
        np.linspace(-120, -500, 10),
        np.linspace(-600, -6000, 20)
    ])

    croco = Model("croco_native")
    vars_to_process = ["temp", "salt", "u", "v"]

    print("\nStarting memory-isolated processing loop...")

    for file_path in tqdm(file_list, desc="Processing CROCO files", unit="file"):
        filename = os.path.basename(file_path)
        output_path = os.path.join(output_dir, f"z_{filename}")

        if os.path.exists(output_path):
            if os.path.getsize(output_path) > 1024:
                continue
            else:
                os.remove(output_path)

        try:
            with dask.config.set(scheduler="threads", num_workers=8):
                ds, xgrid = io.open_files(
                    croco,
                    gridname=grid_master,
                    filenames=[file_path],
                    grid_metrics=2,
                )

                clean_attrs = {}
                for k, v in ds.attrs.items():
                    if isinstance(v, np.ndarray) and v.ndim == 0:
                        clean_attrs[k] = v.item()
                    else:
                        clean_attrs[k] = v

                processed_data = {}

                for var in vars_to_process:
                    if var in ds:
                        da = ds[var]

                        v_dims = [dim for dim in da.dims if dim in ["s", "s_rho", "s_w"]]
                        if v_dims:
                            da = da.chunk({v_dims[0]: -1})

                        res = gop.interp_regular(
                            da=da,
                            grid=xgrid,
                            axis="z",
                            tgrid=tgrid,
                        )

                        if res is not None:
                            processed_data[var] = res.compute()

                if processed_data:
                    ds_out = xr.Dataset(processed_data, attrs=clean_attrs)
                    ds_out.to_netcdf(output_path, engine="netcdf4")
                    ds_out.close()

                ds.close()

        except Exception as e:
            print(f"\n--> Error handling file {filename}: {e}")

            if os.path.exists(output_path):
                try:
                    os.remove(output_path)
                except Exception:
                    pass

    print("\nProcessing complete.")

if __name__ == "__main__":
    main()
