"""Single-parabolic-band transport and Hall properties."""

from __future__ import annotations

from dataclasses import dataclass
from math import exp, log, pi, sqrt

from .band import density_of_states_prefactor, effective_density_of_states
from .constants import (
    BOLTZMANN_CONSTANT as kB,
    ELEMENTARY_CHARGE as e,
    ELECTRON_MASS as me,
    PLANCK_CONSTANT as h,
)
from .fermi import FermiIntegral_fast
from .mobility import ScatteringModel
from .results import HallProperties, TransportProperties

_K_F12 = 2.0 / sqrt(pi)
_K_SIGMA0 = 16.0 * pi * e**2 * me / (3.0 * h**3)
_K_SEEBECK_DEGENERATE = (
    8.0
    * pi**2
    * kB**2
    / (3.0 * e * h**2)
    * (pi / 3.0) ** (2.0 / 3.0)
)


def _validate_positive(name: str, value: float) -> float:
    value = float(value)
    if value <= 0.0:
        raise ValueError(f"{name} must be positive, got {value!r}")
    return value


def _validate_finite_nonzero(name: str, value: float) -> float:
    value = float(value)
    if value == 0.0:
        raise ZeroDivisionError(f"{name} became zero in the transport calculation")
    return value


@dataclass(frozen=True, slots=True)
class SingleParabolicBand:
    """Single isotropic parabolic band.

    Parameters
    ----------
    effective_mass:
        Density-of-states effective mass in units of the free-electron mass.
    band_edge_ev:
        Absolute conduction- or valence-band edge in eV.  The core methods use
        the carrier reduced chemical potential ``eta``: ``(EF-EC)/kBT`` for
        electrons and ``(EV-EF)/kBT`` for holes.
    """

    effective_mass: float
    band_edge_ev: float = 0.0

    def __post_init__(self) -> None:
        _validate_positive("effective_mass", self.effective_mass)

    def effective_density_of_states(self, temperature: float) -> float:
        return effective_density_of_states(self.effective_mass, temperature)

    def density_of_states_prefactor(self) -> float:
        return density_of_states_prefactor(self.effective_mass)

    @staticmethod
    def reduced_chemical_potential_ev(eta: float, temperature: float) -> float:
        """Return ``eta*kBT`` in eV using the carrier-band convention."""

        temperature = _validate_positive("temperature", temperature)
        return eta * kB * temperature / e

    def eta_from_fermi_level(
        self, fermi_level_ev: float, temperature: float, charge_sign: int
    ) -> float:
        """Convert an absolute Fermi level to carrier ``eta``.

        For electrons (``charge_sign=-1``), eta=(EF-EC)/kBT.
        For holes (``charge_sign=+1``), eta=(EV-EF)/kBT.
        """

        if charge_sign not in (-1, 1):
            raise ValueError("charge_sign must be +1 or -1")
        temperature = _validate_positive("temperature", temperature)
        return (
            -charge_sign
            * (fermi_level_ev - self.band_edge_ev)
            * e
            / (kB * temperature)
        )

    def fermi_level_from_eta(
        self, eta: float, temperature: float, charge_sign: int
    ) -> float:
        """Convert carrier ``eta`` to an absolute Fermi level in eV."""

        if charge_sign not in (-1, 1):
            raise ValueError("charge_sign must be +1 or -1")
        return self.band_edge_ev - charge_sign * self.reduced_chemical_potential_ev(
            eta, temperature
        )

    def carrier_density(self, eta: float, temperature: float) -> float:
        nc = self.effective_density_of_states(temperature)
        return nc * _K_F12 * FermiIntegral_fast(float(eta), 0.5)

    def calculate_transport_properties(
        self,
        eta: float,
        temperature: float,
        scattering: ScatteringModel,
        lattice_thermal_conductivity: float = 0.0,
    ) -> TransportProperties:
        """Calculate SPB transport coefficients for a dimensionless ``eta``.

        Units are encoded in the result field names.  The Seebeck and Hall
        signs are controlled by ``scattering.charge_sign``.
        """

        temperature = _validate_positive("temperature", temperature)
        if lattice_thermal_conductivity < 0.0:
            raise ValueError("lattice_thermal_conductivity must be non-negative")

        r = scattering.scattering_factor
        kbt = kB * temperature
        l0_norm = scattering.normalized_mean_free_path

        f05 = FermiIntegral_fast(eta, 0.5)
        fr = _validate_finite_nonzero("F_r", FermiIntegral_fast(eta, r))
        fr1 = FermiIntegral_fast(eta, r + 1.0)
        fr2 = FermiIntegral_fast(eta, r + 2.0)

        nc = self.effective_density_of_states(temperature)
        density = nc * _K_F12 * f05
        if density <= 0.0:
            raise ArithmeticError("calculated carrier density is not positive")

        conductivity = (
            _K_SIGMA0
            * l0_norm
            * self.effective_mass
            * kbt ** (r + 1.0)
            * (r + 1.0)
            * fr
            / 100.0
        )
        mobility = conductivity / (e * density)
        tau_average_fs = (
            mobility * 1.0e-4 * self.effective_mass * me / e * 1.0e15
        )

        seebeck = (
            (kB / e)
            * ((r + 2.0) / (r + 1.0) * fr1 / fr - eta)
            * 1.0e6
            / scattering.charge_sign
        )

        l1 = (r + 1.0) * (r + 3.0) * fr * fr2
        l2 = ((r + 2.0) * fr1) ** 2
        l3 = ((r + 1.0) * fr) ** 2
        lorenz = (kB / e) ** 2 * (l1 - l2) / l3

        # conductivity is in S/cm; divide by 0.01 to obtain S/m.
        kappa_e = lorenz * (conductivity / 0.01) * temperature
        kappa_total = kappa_e + lattice_thermal_conductivity

        power_factor_w_m_k2 = (seebeck * 1.0e-6) ** 2 * (conductivity / 0.01)
        # 1 W/m = 10^4 uW/cm.
        power_factor_uw_cm_k2 = power_factor_w_m_k2 * 1.0e4
        zt = power_factor_w_m_k2 * temperature / kappa_total

        return TransportProperties(
            eta=float(eta),
            temperature=temperature,
            reduced_chemical_potential_ev=self.reduced_chemical_potential_ev(
                eta, temperature
            ),
            carrier_density_cm3=density,
            conductivity_s_cm=conductivity,
            mobility_cm2_v_s=mobility,
            average_relaxation_time_fs=tau_average_fs,
            seebeck_uv_k=seebeck,
            electronic_thermal_conductivity_w_m_k=kappa_e,
            total_thermal_conductivity_w_m_k=kappa_total,
            lorenz_number_w_ohm_k2=lorenz,
            power_factor_w_m_k2=power_factor_w_m_k2,
            power_factor_uw_cm_k2=power_factor_uw_cm_k2,
            zt=zt,
        )

    def calculate_transport_from_fermi_level(
        self,
        fermi_level_ev: float,
        temperature: float,
        scattering: ScatteringModel,
        lattice_thermal_conductivity: float = 0.0,
    ) -> TransportProperties:
        eta = self.eta_from_fermi_level(
            fermi_level_ev, temperature, scattering.charge_sign
        )
        return self.calculate_transport_properties(
            eta, temperature, scattering, lattice_thermal_conductivity
        )

    def calculate_hall_properties(
        self,
        eta: float,
        temperature: float,
        scattering: ScatteringModel,
    ) -> HallProperties:
        """Calculate single-carrier Hall properties for dimensionless ``eta``."""

        temperature = _validate_positive("temperature", temperature)
        r = scattering.scattering_factor
        if r <= -0.25:
            raise ValueError(
                "Hall calculation requires scattering_factor > -0.25 "
                "because F_(2r-1/2) otherwise diverges at the band edge"
            )

        kbt = kB * temperature
        l0_norm = scattering.normalized_mean_free_path
        f05 = FermiIntegral_fast(eta, 0.5)
        fr = _validate_finite_nonzero("F_r", FermiIntegral_fast(eta, r))
        f2r_half = FermiIntegral_fast(eta, 2.0 * r - 0.5)

        nc = self.effective_density_of_states(temperature)
        density = nc * _K_F12 * f05
        if density <= 0.0:
            raise ArithmeticError("calculated carrier density is not positive")

        conductivity = (
            _K_SIGMA0
            * l0_norm
            * self.effective_mass
            * kbt ** (r + 1.0)
            * (r + 1.0)
            * fr
            / 100.0
        )
        mobility = conductivity / (e * density)
        tau_average_fs = (
            mobility * 1.0e-4 * self.effective_mass * me / e * 1.0e15
        )

        hall_factor = (
            1.5
            * (2.0 * r + 0.5)
            / (r + 1.0) ** 2
            * f05
            * f2r_half
            / fr**2
        )
        classical_rh = 1.0 / (e * density * scattering.charge_sign)
        hall_coefficient = hall_factor * classical_rh
        hall_density = density / hall_factor
        hall_mobility = hall_factor * mobility

        return HallProperties(
            eta=float(eta),
            temperature=temperature,
            reduced_chemical_potential_ev=self.reduced_chemical_potential_ev(
                eta, temperature
            ),
            carrier_density_cm3=density,
            conductivity_s_cm=conductivity,
            drift_mobility_cm2_v_s=mobility,
            average_relaxation_time_fs=tau_average_fs,
            hall_factor=hall_factor,
            hall_coefficient_cm3_c=hall_coefficient,
            classical_hall_coefficient_cm3_c=classical_rh,
            hall_density_cm3=hall_density,
            hall_mobility_cm2_v_s=hall_mobility,
        )

    def calculate_hall_from_fermi_level(
        self,
        fermi_level_ev: float,
        temperature: float,
        scattering: ScatteringModel,
    ) -> HallProperties:
        eta = self.eta_from_fermi_level(
            fermi_level_ev, temperature, scattering.charge_sign
        )
        return self.calculate_hall_properties(eta, temperature, scattering)

    def seebeck_nondegenerate_uv_k(
        self,
        carrier_density_cm3: float,
        temperature: float,
        scattering: ScatteringModel,
    ) -> float:
        """Non-degenerate Seebeck approximation in microvolt/K."""

        density = _validate_positive("carrier_density_cm3", carrier_density_cm3)
        nc = self.effective_density_of_states(temperature)
        return (
            kB
            / e
            * (log(nc / density) + scattering.scattering_factor + 2.0)
            * 1.0e6
            / scattering.charge_sign
        )

    def seebeck_degenerate_uv_k(
        self,
        carrier_density_cm3: float,
        temperature: float,
        scattering: ScatteringModel,
    ) -> float:
        """Degenerate Seebeck approximation in microvolt/K."""

        density = _validate_positive("carrier_density_cm3", carrier_density_cm3)
        temperature = _validate_positive("temperature", temperature)
        return (
            _K_SEEBECK_DEGENERATE
            * self.effective_mass
            * me
            * temperature
            * (density * 1.0e6) ** (-2.0 / 3.0)
            * (scattering.scattering_factor + 1.0)
            * 1.0e6
            / scattering.charge_sign
        )

    @staticmethod
    def approximate_lorenz_number_from_seebeck(
        seebeck_uv_k: float,
    ) -> float:
        """Empirical Lorenz-number approximation in W ohm K^-2.

        This preserves the approximation used by the original program.
        The magnitude is used so that electron and hole inputs give the same
        Lorenz estimate.
        """

        return (1.5 + exp(-abs(seebeck_uv_k) / 116.0)) * 1.0e-8


__all__ = ["SingleParabolicBand"]
