"""Convenience functions for eta sweeps."""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np

from .mobility import ScatteringModel
from .results import HallProperties, TransportProperties
from .single_carrier import SingleParabolicBand


@dataclass(frozen=True, slots=True)
class TransportSweep:
    eta: np.ndarray
    reduced_chemical_potential_ev: np.ndarray
    carrier_density_cm3: np.ndarray
    conductivity_s_cm: np.ndarray
    mobility_cm2_v_s: np.ndarray
    relaxation_time_at_fermi_fs: np.ndarray
    average_relaxation_time_fs: np.ndarray
    seebeck_uv_k: np.ndarray
    seebeck_nondegenerate_uv_k: np.ndarray
    seebeck_degenerate_uv_k: np.ndarray
    electronic_thermal_conductivity_w_m_k: np.ndarray
    total_thermal_conductivity_w_m_k: np.ndarray
    lorenz_number_w_ohm_k2: np.ndarray
    approximate_lorenz_number_w_ohm_k2: np.ndarray
    power_factor_w_m_k2: np.ndarray
    power_factor_uw_cm_k2: np.ndarray
    zt: np.ndarray


@dataclass(frozen=True, slots=True)
class HallSweep:
    eta: np.ndarray
    reduced_chemical_potential_ev: np.ndarray
    carrier_density_cm3: np.ndarray
    conductivity_s_cm: np.ndarray
    drift_mobility_cm2_v_s: np.ndarray
    average_relaxation_time_fs: np.ndarray
    hall_factor: np.ndarray
    hall_coefficient_cm3_c: np.ndarray
    classical_hall_coefficient_cm3_c: np.ndarray
    hall_density_cm3: np.ndarray
    hall_mobility_cm2_v_s: np.ndarray


def eta_grid(xmin: float, xmax: float, count: int) -> np.ndarray:
    if count < 2:
        raise ValueError("count must be at least 2")
    if xmax <= xmin:
        raise ValueError("xmax must be greater than xmin")
    return np.linspace(float(xmin), float(xmax), int(count))


def calculate_transport_sweep(
    band: SingleParabolicBand,
    scattering: ScatteringModel,
    temperature: float,
    eta_values: np.ndarray,
    lattice_thermal_conductivity: float = 0.0,
) -> TransportSweep:
    eta_values = np.asarray(eta_values, dtype=float)
    results: list[TransportProperties] = [
        band.calculate_transport_properties(
            eta,
            temperature,
            scattering,
            lattice_thermal_conductivity,
        )
        for eta in eta_values
    ]

    ef = np.array([r.reduced_chemical_potential_ev for r in results])
    tau_ef = np.array(
        [
            np.nan
            if energy <= 0.0
            else scattering.relaxation_time(energy, band.effective_mass) * 1.0e15
            for energy in ef
        ],
        dtype=float,
    )
    density = np.array([r.carrier_density_cm3 for r in results])
    seebeck = np.array([r.seebeck_uv_k for r in results])

    return TransportSweep(
        eta=eta_values.copy(),
        reduced_chemical_potential_ev=ef,
        carrier_density_cm3=density,
        conductivity_s_cm=np.array([r.conductivity_s_cm for r in results]),
        mobility_cm2_v_s=np.array([r.mobility_cm2_v_s for r in results]),
        relaxation_time_at_fermi_fs=tau_ef,
        average_relaxation_time_fs=np.array(
            [r.average_relaxation_time_fs for r in results]
        ),
        seebeck_uv_k=seebeck,
        seebeck_nondegenerate_uv_k=np.array(
            [
                band.seebeck_nondegenerate_uv_k(d, temperature, scattering)
                for d in density
            ]
        ),
        seebeck_degenerate_uv_k=np.array(
            [band.seebeck_degenerate_uv_k(d, temperature, scattering) for d in density]
        ),
        electronic_thermal_conductivity_w_m_k=np.array(
            [r.electronic_thermal_conductivity_w_m_k for r in results]
        ),
        total_thermal_conductivity_w_m_k=np.array(
            [r.total_thermal_conductivity_w_m_k for r in results]
        ),
        lorenz_number_w_ohm_k2=np.array(
            [r.lorenz_number_w_ohm_k2 for r in results]
        ),
        approximate_lorenz_number_w_ohm_k2=np.array(
            [band.approximate_lorenz_number_from_seebeck(s) for s in seebeck]
        ),
        power_factor_w_m_k2=np.array([r.power_factor_w_m_k2 for r in results]),
        power_factor_uw_cm_k2=np.array(
            [r.power_factor_uw_cm_k2 for r in results]
        ),
        zt=np.array([r.zt for r in results]),
    )


def calculate_hall_sweep(
    band: SingleParabolicBand,
    scattering: ScatteringModel,
    temperature: float,
    eta_values: np.ndarray,
) -> HallSweep:
    eta_values = np.asarray(eta_values, dtype=float)
    results: list[HallProperties] = [
        band.calculate_hall_properties(eta, temperature, scattering)
        for eta in eta_values
    ]
    return HallSweep(
        eta=eta_values.copy(),
        reduced_chemical_potential_ev=np.array(
            [r.reduced_chemical_potential_ev for r in results]
        ),
        carrier_density_cm3=np.array([r.carrier_density_cm3 for r in results]),
        conductivity_s_cm=np.array([r.conductivity_s_cm for r in results]),
        drift_mobility_cm2_v_s=np.array(
            [r.drift_mobility_cm2_v_s for r in results]
        ),
        average_relaxation_time_fs=np.array(
            [r.average_relaxation_time_fs for r in results]
        ),
        hall_factor=np.array([r.hall_factor for r in results]),
        hall_coefficient_cm3_c=np.array(
            [r.hall_coefficient_cm3_c for r in results]
        ),
        classical_hall_coefficient_cm3_c=np.array(
            [r.classical_hall_coefficient_cm3_c for r in results]
        ),
        hall_density_cm3=np.array([r.hall_density_cm3 for r in results]),
        hall_mobility_cm2_v_s=np.array(
            [r.hall_mobility_cm2_v_s for r in results]
        ),
    )


__all__ = [
    "TransportSweep",
    "HallSweep",
    "eta_grid",
    "calculate_transport_sweep",
    "calculate_hall_sweep",
]
