"""Command-line application for SPB transport and Hall sweeps."""

from __future__ import annotations

import argparse
from pathlib import Path
from typing import Iterable

import numpy as np

from .band import density_of_states_prefactor
from .constants import LORENZ_NUMBER_FREE_ELECTRON
from .mobility import ScatteringModel
from .single_carrier import SingleParabolicBand
from .sweep import (
    HallSweep,
    TransportSweep,
    calculate_hall_sweep,
    calculate_transport_sweep,
    eta_grid,
)


def _require_openpyxl():
    try:
        import openpyxl
    except ImportError as exc:  # pragma: no cover - environment dependent
        raise RuntimeError(
            "Excel output requires openpyxl. Install with: pip install 'tktransport[app]'"
        ) from exc
    return openpyxl


def _require_matplotlib():
    try:
        from matplotlib import pyplot as plt
    except ImportError as exc:  # pragma: no cover - environment dependent
        raise RuntimeError(
            "Plotting requires matplotlib. Install with: pip install 'tktransport[app]'"
        ) from exc
    return plt


def _to_builtin(value):
    if isinstance(value, np.generic):
        return value.item()
    return value


def _write_parameter_sheet(ws, parameters: dict[str, object]) -> None:
    ws.title = "parameters"
    ws.append(["parameter", "value"])
    for key, value in parameters.items():
        ws.append([key, _to_builtin(value)])
    ws.freeze_panes = "A2"
    ws.column_dimensions["A"].width = 36
    ws.column_dimensions["B"].width = 24


def _write_data_sheet(ws, columns: list[tuple[str, np.ndarray]]) -> None:
    ws.title = "data"
    ws.append([label for label, _ in columns])
    row_count = max((len(values) for _, values in columns), default=0)
    for row_index in range(row_count):
        ws.append(
            [
                _to_builtin(values[row_index]) if row_index < len(values) else None
                for _, values in columns
            ]
        )
    ws.freeze_panes = "A2"
    ws.auto_filter.ref = ws.dimensions
    for column_cells in ws.columns:
        letter = column_cells[0].column_letter
        ws.column_dimensions[letter].width = min(
            24,
            max(12, max(len(str(cell.value or "")) for cell in column_cells) + 2),
        )


def save_transport_excel(
    path: Path,
    sweep: TransportSweep,
    parameters: dict[str, object],
) -> None:
    openpyxl = _require_openpyxl()
    workbook = openpyxl.Workbook()
    _write_parameter_sheet(workbook.active, parameters)
    ws = workbook.create_sheet("data")
    _write_data_sheet(
        ws,
        [
            ("eta=(EF-Eedge)/kBT", sweep.eta),
            ("eta*kBT (eV)", sweep.reduced_chemical_potential_ev),
            ("N (cm^-3)", sweep.carrier_density_cm3),
            ("sigma (S/cm)", sweep.conductivity_s_cm),
            ("mu (cm^2/V/s)", sweep.mobility_cm2_v_s),
            ("tau(EF) (fs)", sweep.relaxation_time_at_fermi_fs),
            ("<tau> (fs)", sweep.average_relaxation_time_fs),
            ("S (uV/K)", sweep.seebeck_uv_k),
            ("S non-degenerate (uV/K)", sweep.seebeck_nondegenerate_uv_k),
            ("S degenerate (uV/K)", sweep.seebeck_degenerate_uv_k),
            ("kappa_e (W/m/K)", sweep.electronic_thermal_conductivity_w_m_k),
            ("kappa_total (W/m/K)", sweep.total_thermal_conductivity_w_m_k),
            ("Lorenz (W ohm/K^2)", sweep.lorenz_number_w_ohm_k2),
            (
                "Lorenz approximate (W ohm/K^2)",
                sweep.approximate_lorenz_number_w_ohm_k2,
            ),
            ("PF (W/m/K^2)", sweep.power_factor_w_m_k2),
            ("PF (uW/cm/K^2)", sweep.power_factor_uw_cm_k2),
            ("ZT", sweep.zt),
        ],
    )
    path.parent.mkdir(parents=True, exist_ok=True)
    workbook.save(path)


def save_hall_excel(
    path: Path,
    sweep: HallSweep,
    parameters: dict[str, object],
) -> None:
    openpyxl = _require_openpyxl()
    workbook = openpyxl.Workbook()
    _write_parameter_sheet(workbook.active, parameters)
    ws = workbook.create_sheet("data")
    _write_data_sheet(
        ws,
        [
            ("eta=(EF-Eedge)/kBT", sweep.eta),
            ("eta*kBT (eV)", sweep.reduced_chemical_potential_ev),
            ("N (cm^-3)", sweep.carrier_density_cm3),
            ("sigma (S/cm)", sweep.conductivity_s_cm),
            ("mu (cm^2/V/s)", sweep.drift_mobility_cm2_v_s),
            ("<tau> (fs)", sweep.average_relaxation_time_fs),
            ("Hall factor", sweep.hall_factor),
            ("RH (cm^3/C)", sweep.hall_coefficient_cm3_c),
            ("RH0=1/(qN) (cm^3/C)", sweep.classical_hall_coefficient_cm3_c),
            ("N_Hall (cm^-3)", sweep.hall_density_cm3),
            ("mu_Hall (cm^2/V/s)", sweep.hall_mobility_cm2_v_s),
        ],
    )
    path.parent.mkdir(parents=True, exist_ok=True)
    workbook.save(path)


def _finish_figure(fig, figure_path: Path | None, show: bool) -> None:
    plt = _require_matplotlib()
    fig.tight_layout()
    if figure_path is not None:
        figure_path.parent.mkdir(parents=True, exist_ok=True)
        fig.savefig(figure_path, dpi=180, bbox_inches="tight")
        print(f"Saved figure: {figure_path}")
    if show:
        plt.show()
    else:
        plt.close(fig)


def plot_transport(
    sweep: TransportSweep,
    title: str,
    figure_path: Path | None,
    show: bool,
) -> None:
    plt = _require_matplotlib()
    fig = plt.figure(figsize=(13, 9))
    axes = [fig.add_subplot(3, 4, i) for i in range(1, 12)]
    (
        ax_tau,
        ax_n,
        ax_sigma,
        ax_mu,
        ax_s,
        ax_kappa,
        ax_pf,
        ax_zt,
        ax_l_eta,
        ax_l_n,
        ax_l_s,
    ) = axes

    finite_tau = np.isfinite(sweep.relaxation_time_at_fermi_fs)
    if finite_tau.any():
        ax_tau.plot(
            sweep.eta[finite_tau],
            sweep.relaxation_time_at_fermi_fs[finite_tau],
            label=r"$\tau(E_F)$",
        )
    ax_tau.plot(sweep.eta, sweep.average_relaxation_time_fs, label=r"$\langle\tau\rangle$")
    ax_tau.set_xlabel(r"$\eta=(E_F-E_{edge})/k_BT$")
    ax_tau.set_ylabel(r"$\tau$ (fs)")
    ax_tau.legend(fontsize=8)

    ax_n.plot(sweep.eta, sweep.carrier_density_cm3)
    ax_n.set_xlabel(r"$\eta$")
    ax_n.set_ylabel(r"$N$ (cm$^{-3}$)")
    ax_n.set_yscale("log")

    ax_sigma.plot(sweep.carrier_density_cm3, sweep.conductivity_s_cm)
    ax_sigma.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_sigma.set_ylabel(r"$\sigma$ (S/cm)")
    ax_sigma.set_xscale("log")
    ax_sigma.set_yscale("log")
    ax_sigma.set_title(title, fontsize=10)

    ax_mu.plot(sweep.carrier_density_cm3, sweep.mobility_cm2_v_s)
    ax_mu.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_mu.set_ylabel(r"$\mu$ (cm$^2$/V/s)")
    ax_mu.set_xscale("log")

    ax_s.plot(sweep.carrier_density_cm3, sweep.seebeck_uv_k, label="SPB")
    ax_s.plot(
        sweep.carrier_density_cm3,
        sweep.seebeck_nondegenerate_uv_k,
        linestyle="--",
        label="non-degenerate",
    )
    ax_s.plot(
        sweep.carrier_density_cm3,
        sweep.seebeck_degenerate_uv_k,
        linestyle="--",
        label="degenerate",
    )
    ax_s.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_s.set_ylabel(r"$S$ ($\mu$V/K)")
    ax_s.set_xscale("log")
    ax_s.legend(fontsize=7)

    ax_kappa.plot(
        sweep.carrier_density_cm3,
        sweep.electronic_thermal_conductivity_w_m_k,
    )
    ax_kappa.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_kappa.set_ylabel(r"$\kappa_e$ (W/m/K)")
    ax_kappa.set_xscale("log")
    ax_kappa.set_yscale("log")

    ax_pf.plot(sweep.carrier_density_cm3, sweep.power_factor_uw_cm_k2)
    ax_pf.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_pf.set_ylabel(r"PF ($\mu$W/cm/K$^2$)")
    ax_pf.set_xscale("log")

    ax_zt.plot(sweep.carrier_density_cm3, sweep.zt)
    ax_zt.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_zt.set_ylabel("ZT")
    ax_zt.set_xscale("log")

    ax_l_eta.plot(sweep.eta, sweep.lorenz_number_w_ohm_k2, label="calculated")
    ax_l_eta.axhline(
        LORENZ_NUMBER_FREE_ELECTRON,
        linestyle="--",
        label="degenerate limit",
    )
    ax_l_eta.set_xlabel(r"$\eta$")
    ax_l_eta.set_ylabel(r"$L$ (W$\Omega$/K$^2$)")
    ax_l_eta.legend(fontsize=7)

    ax_l_n.plot(
        sweep.carrier_density_cm3,
        sweep.lorenz_number_w_ohm_k2,
        label="calculated",
    )
    ax_l_n.axhline(LORENZ_NUMBER_FREE_ELECTRON, linestyle="--")
    ax_l_n.set_xlabel(r"$N$ (cm$^{-3}$)")
    ax_l_n.set_ylabel(r"$L$ (W$\Omega$/K$^2$)")
    ax_l_n.set_xscale("log")

    ax_l_s.plot(
        sweep.seebeck_uv_k,
        sweep.lorenz_number_w_ohm_k2,
        label="calculated",
    )
    ax_l_s.plot(
        sweep.seebeck_uv_k,
        sweep.approximate_lorenz_number_w_ohm_k2,
        label="approximate",
    )
    ax_l_s.axhline(LORENZ_NUMBER_FREE_ELECTRON, linestyle="--")
    ax_l_s.set_xlabel(r"$S$ ($\mu$V/K)")
    ax_l_s.set_ylabel(r"$L$ (W$\Omega$/K$^2$)")
    ax_l_s.legend(fontsize=7)

    _finish_figure(fig, figure_path, show)


def plot_hall(
    sweep: HallSweep,
    title: str,
    figure_path: Path | None,
    show: bool,
) -> None:
    plt = _require_matplotlib()
    fig = plt.figure(figsize=(12, 8))
    axes = [fig.add_subplot(3, 3, i) for i in range(1, 9)]
    ax_tau, ax_n, ax_sigma, ax_mu, ax_rh, ax_fh, ax_rh_n, ax_fh_n = axes

    ax_tau.plot(sweep.eta, sweep.average_relaxation_time_fs)
    ax_tau.set_xlabel(r"$\eta$")
    ax_tau.set_ylabel(r"$\langle\tau\rangle$ (fs)")

    ax_n.plot(sweep.eta, sweep.carrier_density_cm3, label="drift density")
    ax_n.plot(sweep.eta, sweep.hall_density_cm3, label="Hall density")
    ax_n.set_xlabel(r"$\eta$")
    ax_n.set_ylabel(r"$N$ (cm$^{-3}$)")
    ax_n.set_yscale("log")
    ax_n.legend(fontsize=8)

    ax_sigma.plot(sweep.hall_density_cm3, sweep.conductivity_s_cm)
    ax_sigma.set_xlabel(r"$N_{Hall}$ (cm$^{-3}$)")
    ax_sigma.set_ylabel(r"$\sigma$ (S/cm)")
    ax_sigma.set_xscale("log")
    ax_sigma.set_yscale("log")
    ax_sigma.set_title(title, fontsize=10)

    ax_mu.plot(
        sweep.hall_density_cm3,
        sweep.drift_mobility_cm2_v_s,
        label="drift mobility",
    )
    ax_mu.plot(
        sweep.hall_density_cm3,
        sweep.hall_mobility_cm2_v_s,
        label="Hall mobility",
    )
    ax_mu.set_xlabel(r"$N_{Hall}$ (cm$^{-3}$)")
    ax_mu.set_ylabel(r"$\mu$ (cm$^2$/V/s)")
    ax_mu.set_xscale("log")
    ax_mu.legend(fontsize=8)

    ax_rh.plot(
        sweep.hall_density_cm3,
        np.abs(sweep.classical_hall_coefficient_cm3_c),
        label=r"$|R_{H0}|$",
    )
    ax_rh.plot(
        sweep.hall_density_cm3,
        np.abs(sweep.hall_coefficient_cm3_c),
        label=r"$|R_H|$",
    )
    ax_rh.set_xlabel(r"$N_{Hall}$ (cm$^{-3}$)")
    ax_rh.set_ylabel(r"$|R_H|$ (cm$^3$/C)")
    ax_rh.set_xscale("log")
    ax_rh.set_yscale("log")
    ax_rh.legend(fontsize=8)

    ax_fh.plot(sweep.hall_density_cm3, sweep.hall_factor)
    ax_fh.set_xlabel(r"$N_{Hall}$ (cm$^{-3}$)")
    ax_fh.set_ylabel(r"$F_H$")
    ax_fh.set_xscale("log")

    ax_rh_n.plot(sweep.hall_density_cm3, np.abs(sweep.hall_coefficient_cm3_c))
    ax_rh_n.set_xlabel(r"$N_{Hall}$ (cm$^{-3}$)")
    ax_rh_n.set_ylabel(r"$|R_H|$ (cm$^3$/C)")
    ax_rh_n.set_xscale("log")
    ax_rh_n.set_yscale("log")

    ax_fh_n.plot(sweep.hall_density_cm3, sweep.hall_factor)
    ax_fh_n.set_xlabel(r"$N_{Hall}$ (cm$^{-3}$)")
    ax_fh_n.set_ylabel(r"$F_H$")
    ax_fh_n.set_xscale("log")

    _finish_figure(fig, figure_path, show)


def _print_rows(headers: Iterable[str], rows: Iterable[Iterable[float]]) -> None:
    print(" ".join(f"{header:>14}" for header in headers))
    for row in rows:
        print(" ".join(f"{value:14.6g}" for value in row))


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(
        description="Single-parabolic-band transport / Hall calculation"
    )
    parser.add_argument("mode", choices=("prop", "hall"))
    parser.add_argument("--temperature", "-T", type=float, default=300.0, help="K")
    parser.add_argument(
        "--effective-mass", "-m", type=float, default=0.3, help="in electron masses"
    )
    parser.add_argument(
        "--scattering-factor", "-r", type=float, default=0.5, help="power-law r"
    )
    parser.add_argument(
        "--mean-free-path", "-l", type=float, default=1.0e-8, help="l0 in m"
    )
    parser.add_argument(
        "--lattice-thermal-conductivity",
        "-k",
        type=float,
        default=5.0,
        help="W/m/K",
    )
    parser.add_argument("--carrier", choices=("electron", "hole"), default="hole")
    parser.add_argument("--xmin", type=float, default=-30.0, help="minimum eta")
    parser.add_argument("--xmax", type=float, default=250.0, help="maximum eta")
    parser.add_argument("--nx", type=int, default=281, help="number of eta points")
    parser.add_argument("-o", "--output", type=Path, default=None, help="output xlsx")
    parser.add_argument("--figure", type=Path, default=None, help="save plot image")
    parser.add_argument("--no-show", action="store_true", help="do not open plot window")
    parser.add_argument(
        "--no-plot", action="store_true", help="skip plotting entirely"
    )
    parser.add_argument(
        "--print-every",
        type=int,
        default=10,
        help="print every Nth row; 0 disables table output",
    )
    return parser


def main(argv: list[str] | None = None) -> int:
    args = build_parser().parse_args(argv)
    if args.no_plot and args.figure is not None:
        raise ValueError("--figure cannot be used together with --no-plot")
    charge_sign = -1 if args.carrier == "electron" else 1
    band = SingleParabolicBand(args.effective_mass)
    scattering = ScatteringModel(
        charge_sign=charge_sign,
        scattering_factor=args.scattering_factor,
        mean_free_path=args.mean_free_path,
    )
    etas = eta_grid(args.xmin, args.xmax, args.nx)
    output = args.output or Path(f"properties-{args.mode}.xlsx")

    parameters = {
        "mode": args.mode,
        "temperature (K)": args.temperature,
        "effective mass (m_e)": args.effective_mass,
        "carrier": args.carrier,
        "charge sign": charge_sign,
        "scattering factor r": args.scattering_factor,
        "mean free path l0 (m)": args.mean_free_path,
        "lattice thermal conductivity (W/m/K)": args.lattice_thermal_conductivity,
        "eta minimum": args.xmin,
        "eta maximum": args.xmax,
        "eta points": args.nx,
        "effective DOS (cm^-3)": band.effective_density_of_states(args.temperature),
        "DOS prefactor (cm^-3 eV^-3/2)": density_of_states_prefactor(
            args.effective_mass
        ),
        "free-electron Lorenz number (W ohm/K^2)": LORENZ_NUMBER_FREE_ELECTRON,
        "Fermi integral": "tktransport.fermi.FermiIntegral_fast",
    }

    title = (
        f"T={args.temperature:g} K, m*={args.effective_mass:g} me, "
        f"r={args.scattering_factor:g}, carrier={args.carrier}"
    )
    print(title)
    print(f"Effective DOS: {parameters['effective DOS (cm^-3)']:.6g} cm^-3")

    if args.mode == "prop":
        sweep = calculate_transport_sweep(
            band,
            scattering,
            args.temperature,
            etas,
            args.lattice_thermal_conductivity,
        )
        save_transport_excel(output, sweep, parameters)
        if args.print_every > 0:
            indices = range(0, args.nx, args.print_every)
            _print_rows(
                ("eta", "eta*kBT(eV)", "N(cm^-3)", "sigma", "mu", "S", "L", "PF", "ZT"),
                (
                    (
                        sweep.eta[i],
                        sweep.reduced_chemical_potential_ev[i],
                        sweep.carrier_density_cm3[i],
                        sweep.conductivity_s_cm[i],
                        sweep.mobility_cm2_v_s[i],
                        sweep.seebeck_uv_k[i],
                        sweep.lorenz_number_w_ohm_k2[i],
                        sweep.power_factor_uw_cm_k2[i],
                        sweep.zt[i],
                    )
                    for i in indices
                ),
            )
        if not args.no_plot:
            plot_transport(sweep, title, args.figure, not args.no_show)
    else:
        sweep = calculate_hall_sweep(band, scattering, args.temperature, etas)
        save_hall_excel(output, sweep, parameters)
        if args.print_every > 0:
            indices = range(0, args.nx, args.print_every)
            _print_rows(
                ("eta", "eta*kBT(eV)", "N(cm^-3)", "sigma", "mu", "FH", "RH", "N_Hall"),
                (
                    (
                        sweep.eta[i],
                        sweep.reduced_chemical_potential_ev[i],
                        sweep.carrier_density_cm3[i],
                        sweep.conductivity_s_cm[i],
                        sweep.drift_mobility_cm2_v_s[i],
                        sweep.hall_factor[i],
                        sweep.hall_coefficient_cm3_c[i],
                        sweep.hall_density_cm3[i],
                    )
                    for i in indices
                ),
            )
        if not args.no_plot:
            plot_hall(sweep, title, args.figure, not args.no_show)

    print(f"Saved data: {output}")
    return 0


if __name__ == "__main__":  # pragma: no cover
    raise SystemExit(main())
