#!/usr/bin/env python3

from pathlib import Path
import sys

import numpy as np
import pandas as pd


NARROW = [
    "N395", "N419", "N501", "N540", "N662",
    "N673", "N708", "N964", "N1008",
]

MEDIUM = [
    "M411", "M438", "M464", "M490", "M517",
]


def to_angstrom(wave_nm):
    return np.asarray(wave_nm, dtype=float) * 10.0


def to_fraction(transmission, label):
    transmission = np.asarray(transmission, dtype=float)

    label = str(label).lower()

    if "%" in label:
        return transmission / 100.0

    if np.nanmax(transmission) > 1.5:
        return transmission / 100.0

    return transmission


def write_filter(sheet_name, df, outdir):
    df = df.dropna(how="all").dropna(axis=1, how="all")

    if df.shape[1] < 2:
        print(f"skip {sheet_name}: fewer than two columns")
        return False

    wave_col = df.columns[0]
    trans_col = df.columns[1]

    wavelength = to_angstrom(df[wave_col])
    transmission = to_fraction(df[trans_col], trans_col)

    good = np.isfinite(wavelength) & np.isfinite(transmission)
    wavelength = wavelength[good]
    transmission = transmission[good]

    order = np.argsort(wavelength)
    wavelength = wavelength[order]
    transmission = transmission[order]

    outpath = outdir / f"{sheet_name}.dat"

    np.savetxt(
        outpath,
        np.column_stack([wavelength, transmission]),
        delimiter=",",
        header="Wavelength,Transmission",
        comments="",
        fmt=["%.8f", "%.10g"],
    )

    print(
        f"wrote {outpath} "
        f"({wavelength.min():.1f}-{wavelength.max():.1f} A, "
        f"max T={transmission.max():.4g})"
    )

    return True


def write_index(outdir, index_name, names):
    files = [f"{name}.dat" for name in names if (outdir / f"{name}.dat").exists()]
    (outdir / index_name).write_text("\n".join(files) + "\n", encoding="utf-8")


def main():
    if len(sys.argv) != 2:
        raise SystemExit("usage: ./make_decam_medium_narrow.py decam_medium_narrow.xls")

    xls = Path(sys.argv[1]).expanduser().resolve()

    narrow_dir = Path("DECam/Narrow")
    medium_dir = Path("DECam/Medium")

    narrow_dir.mkdir(parents=True, exist_ok=True)
    medium_dir.mkdir(parents=True, exist_ok=True)

    sheets = pd.read_excel(xls, sheet_name=None)

    found_narrow = []
    found_medium = []

    for sheet_name, df in sheets.items():
        name = str(sheet_name).strip()

        if name in NARROW:
            if write_filter(name, df, narrow_dir):
                found_narrow.append(name)

        elif name in MEDIUM:
            if write_filter(name, df, medium_dir):
                found_medium.append(name)

        else:
            print(f"skip {name}: not narrow/medium")

    write_index(narrow_dir, "Narrow", NARROW)
    write_index(medium_dir, "Medium", MEDIUM)

    print()
    print("Narrow found:", " ".join(found_narrow))
    print("Medium found:", " ".join(found_medium))

    missing_narrow = sorted(set(NARROW) - set(found_narrow))
    missing_medium = sorted(set(MEDIUM) - set(found_medium))

    if missing_narrow:
        print("Narrow missing:", " ".join(missing_narrow))

    if missing_medium:
        print("Medium missing:", " ".join(missing_medium))


if __name__ == "__main__":
    main()
