import xarray as xr
from pathlib import Path
import numpy as np

SKIP_DIRS = {"old_data", "new_data", "__pycache__"}  # folder names to skip

def cast_uint8_to_int(ds):
    """Convert all uint8 variables in a dataset to int16, adjusting _FillValue."""
    new_vars = {}
    for var_name, da in ds.data_vars.items():
        if da.dtype == np.uint8:
            fill = da.attrs.get("_FillValue", None)
            new_dtype = np.int16
            data_casted = da.astype(new_dtype)
            if fill is not None:
                if fill > np.iinfo(new_dtype).max:
                    print(f"⚠️ Skipping {var_name}: _FillValue too large for int16")
                    continue
                data_casted.attrs["_FillValue"] = np.int16(fill)
            new_vars[var_name] = data_casted
        else:
            new_vars[var_name] = da
    return xr.Dataset(new_vars, coords=ds.coords, attrs=ds.attrs)


def convert_to_netcdf3_64bit(file_path):
    try:
        ds = xr.open_dataset(file_path, decode_cf=False, chunks = {})
        ds = cast_uint8_to_int(ds)
        tmp_path = file_path.with_suffix(".tmp.nc")
        ds.to_netcdf(tmp_path, format="NETCDF3_64BIT_OFFSET")
        ds.close()
        # tmp_path.replace(file_path)
        print(f"✅ Converted: {file_path}")
    except Exception as e:
        print(f"❌ Failed: {file_path} — {e}")



def should_skip(path):
    return any(skip in path.parts for skip in SKIP_DIRS)


def main():
    base_dir = Path(__file__).parent.resolve()
    for file_path in base_dir.rglob("*.nc"):
        if should_skip(file_path):
            continue
        file_path = Path("/glade/campaign/cesm/cesmdata/cseg/inputdata/ocn/mom/croc/chl/data/LandWater15ARC.nc")
        convert_to_netcdf3_64bit(file_path)
        break


if __name__ == "__main__":
    main()
