from functools import lru_cache
from math import factorial, fsum, sqrt
from typing import Any
from typing import Tuple, Iterable

import numpy as np
from lmfit import Model
from lmfit.lineshapes import linear
from lmfit.models import LinearModel as LmfitLinearModel, update_param_vals
from scipy.optimize import fsolve
from scipy.special import gamma, hyp2f1, jv


##################################################################
# Cosine Model
##################################################################
def fn_cosine(x, baseline, amplitude, frequency, x0=0.0):
    """
    Cosine lineshape

    :param numpy.ndarray x: Independent variable
    :param float baseline: Baseline of cosine
    :param float amplitude: Amplitude of cosine
    :param float frequency: Frequency of oscillation
    :param float x0: starting point in x
    :return: Cosine function
    :rtype: numpy.ndarray
    """
    return baseline + amplitude * np.cos(2 * np.pi * frequency * (x - x0))


class CosineModel(Model):
    __doc__ = """Class for Cosine model"""

    def __init__(self, *args, **kwargs):
        super(CosineModel, self).__init__(fn_cosine, *args, **kwargs)

    __model__ = Model(fn_cosine)

    def guess(self, x, y, **kwargs):
        """
        Takes x and y data and generates initial guesses for
        fit parameters.

        :param numpy.ndarray x: Independent variable X
        :param numpy.ndarray y: Dependent variable Y
        :param kwargs: lmfit parameters
        :return: Fit parameters
        """
        b_guess = np.median(y)
        a_guess = (np.max(y) - np.min(y)) / 2.
        f_guess = get_freq_from_fft(x, y)
        x_guess = x[0]

        pars = self.make_params(baseline=b_guess, amplitude=a_guess, frequency=f_guess)
        pars['amplitude'].value = kwargs.pop('amplitude_guess', a_guess)
        pars['baseline'].value = kwargs.pop('baseline_guess', b_guess)
        pars['frequency'].value = kwargs.pop('detuning_guess', f_guess)
        pars['x0'].value = kwargs.pop('x0_guess', x_guess)
        return update_param_vals(pars, self.prefix, **kwargs)

    def do_fit(self, x, y, errs=None, **kwargs):
        """
        Performs a fit

        :param numpy.ndarray x: Independent variable X
        :param numpy.ndarray y: Dependent variable Y
        :param numpy.ndarray errs: Errors in dependent variable to weight by
        :param kwargs: lmfit parameters
        :return: fit
        """
        par = self.guess(x, y, **kwargs)
        if errs is not None:
            fit = self.fit(y, x=x, weights=1.0 / errs, params=par, nan_policy="omit")
        else:
            fit = self.fit(y, x=x, params=par, nan_policy="omit")
        return fit

    def report_fit(self, x, y, **kwargs):
        """
        Reports the results of a fit.

        :param numpy.ndarray x: Independent variable X
        :param numpy.ndarray y: Dependent variable Y
        :param kwargs: lmfit parameters
        :return: fit
        """

        if not len(x) == len(y):
            raise ValueError("Lengths of x and y arrays must be equal.")

        fit = self.do_fit(x, y, **kwargs)

        return fit


def get_freq_from_fft(x, y, offset=0.0):
    """Returns the frequency at which an fft is peaked in amplitude.
    Notes:
    - A constant offset is subtracted
    - The two largest frequency components are taken in case f=0 is still the
    largest components
    - A sine wave will have peaks at +/- the true frequency. This function
    returns the absolute value.
    """
    assert len(x) > 1
    f = np.fft.fftfreq(len(y), x[1] - x[0])
    A = np.abs(np.fft.fft(np.asarray(y) - offset))
    f_max_2 = f[np.argpartition(A, -2)][-2:]  # f's for largest 2 elements of A
    if f_max_2[0] != 0.0:
        f_max = f_max_2[0]
    else:
        f_max = f_max_2[1]
    return np.abs(f_max)


##################################################################
# Tunable Transmon Functions
##################################################################
def fn_cavity_and_tunable_transmon(x, fr_0, fb_offset, fb_period, ej_ec, ratio, ec, g):
    """
    Model hamiltonian of a tunable transmon and resonator

    :param ndarray x: Biases in Volts.
    :param float fr_0: The resonator frequency in Hz.
    :param float fb_offset: A flux bias offset from 0 associated with an extrema of qubit frequency.
    :param float fb_period: The flux bias period.
    :param float ej_ec: The product of the sum Ej1+Ej2 and Ec.
    :param float ratio: The ratio of Ej1 and Ej2.
    :param float ec: The transmon charging energy.
    :param float g: The cavity-resonator coupling.
    :return: Frequency of the cavity at a given qubit bias.
    :rtype: ndarray
    """
    fq_0 = fn_tunable_transmon_freq(x, fb_offset, fb_period, ej_ec, ratio, ec)

    return fr_0 + g ** 2 / (fr_0 - fq_0)


def fn_tunable_transmon_freq(x, fb_offset, fb_period, ej_ec, ratio, ec):
    """
    Model function for tunable transmon frequency versus applied flux bias. The units of x,
    fb_offset, fb_period, can be either a voltage or a current but must be consistent.

    The Hamiltonian parameters $E_J$ and $E_C$ appear in this model only via the product $E_J E_C$
    and therefore this product is the model parameter which is allowed to vary and whose uncertainty
    can be derived directly from the $\chi^2$ surface. The $E_J$ and $E_C$ can then be derived from
    the product in combination with the anharmonicity.

    It is anticipated that the model will serve best if the anharmonicity is fixed to an
    independently measured value during fitting.

    See eqn (6) in Dephasing under parametric modulation in tunable transmon qubits, Eyob Sete in
    the Papyrus repository for the definition of the qubit transition frequency.

    :param float x: Independent variable of the model. Flux bias value at which to evaluate function
                    [Volts or Amperes].
    :param float fb_offset: Flux-bias offset, the voltage or current where fmax is realized
                            [Volts or Amperes].
    :param float fb_period: Flux-bias period, the voltage or current interval over which a flux
                            quantum is traversed [Volts or Amperes].
    :param float ej_ec: Product EJ * EC, where EJ is the total Josephson energy of both junctions
                        and EC is the charging energy [Hz^2].
    :param float ratio: Ratio of Josephson energies of the two Josephson junctions [Number].
    :param float ec: Transmon charge energy, approximately equal to (f01 - f12) at f01=fmax [Hz].
    :return: Tunable transmon transition frequency f01 [Hz]
    :rtype: float
    """
    flux = (x - fb_offset) / fb_period
    ej_ec_eff = (ej_ec * np.sqrt(1.0 + ratio ** 2.0 + 2.0 * ratio * np.cos(2.0 * np.pi * flux)) /
                 (1.0 + ratio))
    xi = np.sqrt(2 * ec ** 2 / ej_ec)
    return np.sqrt(8.0 * ej_ec_eff) - ec * (1 + xi / 4 + 21 * xi ** 2 / 128 + 19 * xi ** 3 / 128)


##################################################################
# Linear Model
##################################################################
class LinearModel(LmfitLinearModel):
    __doc__ = """Class for Linear model"""

    def __init__(self, *args, **kwargs):
        super(LinearModel, self).__init__(*args, **kwargs)

    __model__ = Model(linear)

    def report_fit(self, x, y, errs=None, **kwargs):
        """Reports the results of a fit. May change depending on what we
        want this to return.
        """

        if not len(x) == len(y):
            raise ValueError("Lengths of x and y arrays must be equal.")
        if not len(x) > 2:
            raise ValueError("You must provide more than 2 data points.")

        par = self.guess(y, x=x, **kwargs)

        if self.error_bars_are_valid(errs):
            fit = self.fit(y, x=x, weights=1.0 / np.array(errs), params=par)
        else:
            fit = self.fit(y, x=x, params=par)

        return fit, (
            fit.params['slope'],
            fit.params['intercept']
        ), True

    def error_bars_are_valid(self, errs: Any) -> bool:
        """Validity check for error bars, before they are used in fitting.

        :param errs: A variable of any type. Typically None or numpy.ndarray.
        :return: True for an ndarray full of non-zero values, false for everything else.
        """
        if isinstance(errs, np.ndarray):
            if not any(errs == 0.0):
                return True
            else:
                print("Invalid value of 0.0 found in fit error bars. All error bars may "
                      "be dropped for fitting.")
                return False
        else:
            return False


##################################################################
# AC-flux-modulated Tunable Transmon Functions & Analysis
##################################################################
def fn_modulation_detuning(vs, fmax, eta_max, fmin, scale, dc_flux=0):
    """
    Model for modulation detuning scan

    :param vs: amplitudes in volts or amps
    :param fmax: qubit f01 frequency at max in Hz
    :param eta_max: qubit anharmonicity at max in Hz
    :param fmin: qubit f01 frequency at min in Hz
    :param scale: 1/the voltage that will modulate you one flux-period's worth
    :param dc_flux: parking flux in phi_0
    :return: modulation_detuning - f01 (if at fmax) (freq_with_mod - freq_without_mod)
    """
    ec, ej1, ej2 = transmon_ec_ej(fmax, fmin, eta_max)
    return [frequency_shift_01(ec, ej1, ej2, dc_flux=dc_flux, ac_flux=v * scale) for v in vs]


# Perturbation theory series
ORDERMAX = 10  # maximal order in perturbation theory expansion. Can go up to 30
PT_ORDERS = range(-1, ORDERMAX + 1)  # the exponents in the perturbation theory expansion

# list of rational numbers f^(p) for f01, f12, anh, written in the compact form (a_p, b_p)
F01_COEF_PARAMS = ((4, 0), (-1, 0), (-1, 2), (-21, 7), (-19, 7), (-5319, 15), (-6649, 15),
                   (-1180581, 22), (-446287, 20), (-1489138635, 31), (-648381403, 29),
                   (-614557854099, 38),)
ANH_COEF_PARAMS = ((0, 0), (1, 0), (9, 4), (81, 7), (3645, 12), (46899, 15), (1329129, 19),
                   (20321361, 22), (2648273373, 28), (45579861135, 31), (1647988255539, 35),
                   (31160327412879, 38),)

coef_gen = lambda x: x[0] / 2 ** x[1]
F01_COEF = tuple(map(coef_gen, F01_COEF_PARAMS))
ANH_COEF = tuple(map(coef_gen, ANH_COEF_PARAMS))

XI_MAX_GUESS = 0.18


def transmon_ec_ej(fmax: float, fmin: float, eta_max: float) -> Tuple[float, float, float]:
    """
    Get ec, ej1, and ej2 for a tunable transmon with a given fmax, fmin, and eta_max.

    :param fmax: Qubit f01 frequency at 0-flux.
    :param fmin: Qubit f01 frequency at half-flux.
    :param eta_max: qubit anharmonicity at 0-flux.
    :return: The Hamiltonian parameters of the transmon.
    """
    # Determine xi_max
    xi_max = get_xi_max(fmax, eta_max)

    # Determine ec
    ec = fmax / f01_xi(xi_max)

    # Determine xi_min
    opt_func_xi_min = lambda xi: f01_xi(xi) - fmin / ec
    xi_min = fsolve(opt_func_xi_min, XI_MAX_GUESS)[0]

    # calculate ej1 and ej2 from ec and xi
    ej1 = ec * (1 / xi_max ** 2 + 1 / xi_min ** 2)
    ej2 = ec * (1 / xi_max ** 2 - 1 / xi_min ** 2)

    return ec, ej1, ej2


def phi_from_freqshift_01(ec: float, ej1: float, ej2: float,
                          dc_flux: float, delta_f01: float, phi_guess: float = 0.25) -> np.float64:
    """
    Find the smallest modulation amplitude in Φ₀ that generates the desired shift in frequency

    :param ec: Charging energy.
    :param ej1: Josephson energy of the first junction.
    :param ej2: Josephson energy of the second junction.
    :param dc_flux: DC flux in Φ₀.
    :param delta_f01: the frequency shift between the average frequency under modulation,
                      and the frequency without modulation: freq_with_mod - freq_without_mod
    :param phi_guess: A guess for the acheived phi in units of Φ₀
    :return: the modulation amplitude in Φ₀ that returns the given frequency shift
    """
    func = lambda ac_flux: abs(frequency_shift_01(ec, ej1, ej2, dc_flux, ac_flux) - delta_f01)
    return fsolve(func, phi_guess)[0]


def frequency_shift_01(ec, ej1, ej2, dc_flux, ac_flux):
    """
    Calculate the shift of the average frequency under modulation at the given amplitude

    :param ec: Charging energy.
    :param ej1: Josephson energy of the first junction.
    :param ej2: Josephson energy of the second junction.
    :param dc_flux: DC flux in Φ₀.
    :param ac_flux: AC flux amplitude in Φ₀.
    :return: Frequency shift of the tunable transmon for the given dc and ac flux as
             freq_with_mod - freq_without_mod
    """
    pt_coeff = f01_coef_fourier(ec, ej1, ej2)

    return (modulated_coef_fourier(dc_flux, ac_flux, pt_coeff, 0) -
            modulated_coef_fourier(dc_flux, 0, pt_coeff, 0))


def modulated_coef_fourier(dc_flux: float, ac_flux: float, pt_coef: Iterable, k: int) -> float:
    """
    Calculate the Fourier coefficient f_k from the list pt_coeff.

    :param dc_flux: DC flux in Φ₀.
    :param ac_flux: AC flux amplitude in Φ₀.
    :param pt_coef: perturbation theory coefficients for the desired transmon and desired transition
     frequency
    :param k: the harmonic of the Fourier expansion.
                 k=0 is the average frequency
                 k=2 is oscillations about the average frequency
                 At fmax and fmin, even ks are the only nonzero harmonics
    :return: the kth fourier coefficient of the transition frequency
    """
    fk = fsum([c * np.cos(n * 2 * np.pi * dc_flux + k * np.pi / 2) *
               jv(k, n * 2 * np.pi * ac_flux) for n, c in enumerate(pt_coef)])

    if k > 0:
        fk *= 2

    return fk


def f01_coef_fourier(ec: float, ej1: float, ej2: float = 0.0) -> Tuple[float]:
    """
    Fourier coefficients for the f01 transition under modulation.

    :param ec: Charging energy.
    :param ej1: Josephson energy of the first junction.
    :param ej2: Josephson energy of the second junction.
    :return: A list of Fourier coefficients.
    """
    return coef_fourier_gen(ec, ej1, ej2, F01_COEF)


def coef_fourier_gen(ec: float, ej1: float, ej2: float, coef: Iterable[Tuple[float, float]]
                     ) -> Tuple[float]:
    """
    Calculate terms for the Fourier expansion on Perturbation theory of Transmons,
    from ec, ej1, ej2, and the series of transmon coefficients.

    :param ec: Charging energy.
    :param ej1: Josephson energy of the first junction.
    :param ej2: Josephson energy of the second junction.
    :param coef: Coefficients for perturbation theory solution to flux-modulated transmon.
    :return: The list of Fourier coefficients.
    """
    assert ej1 > ej2

    xi_bar = sqrt(2 * ec / sqrt(ej1 ** 2 + ej2 ** 2))
    ej_red = 2 * ej1 * ej2 / (ej1 ** 2 + ej2 ** 2)

    s0 = ec * fsum([c * xi_bar ** p * hyp2f1(p / 8, p / 8 + 1 / 2, 1, ej_red ** 2)
                    for c, p in zip(coef, PT_ORDERS)])

    hyp_help = lambda n, p: hyp2f1(n / 2 + p / 8, n / 2 + p / 8 + 1 / 2, n + 1, ej_red ** 2)

    s = [s0]
    for n in range(1, ORDERMAX + 1):
        s_partial = fsum([c * xi_bar ** p * gamma(n + p / 4) / gamma(p / 4) * hyp_help(n, p)
                          for c, p in zip(coef, PT_ORDERS) if p != 0])

        sn = ec * 2 * (-ej_red / 2) ** n / factorial(n) * s_partial
        s.append(sn)

    return tuple(s)


@lru_cache()
def get_xi_max(fmax: float, eta_max: float) -> float:
    """
    Obtain Xi at 0-flux.

    :param fmax: The qubit f01 frequency at 0-flux.
    :param eta_max: The anharmonicity at 0-flux.
    :return: Xi at 0-flux
    """
    opt_func_xi_max = lambda xi: f01_xi(xi) / anh_xi(xi) - fmax / abs(eta_max)
    xi_max = fsolve(opt_func_xi_max, XI_MAX_GUESS)[0]

    return xi_max


def f01_xi(xi: float) -> float:
    """
    Get F01 in units of Ec.

    :param xi: sqrt(2 * ec / ej_eff).
    :return: f01 in units of ec.
    """
    return fsum([coef * xi ** p for coef, p in zip(F01_COEF, PT_ORDERS)])


def anh_xi(xi: float) -> float:
    """
    Get anharmonicity in units of Ec.

    :param xi: sqrt(2 * ec / ej_eff).
    :return: Anharmonicity in units of ec.
    """
    return fsum([coef * xi ** p for coef, p in zip(ANH_COEF, PT_ORDERS)])
