"""Computes the linear response (Hall conductivity) of the Haldane model over
a range of chemical potentials
"""

import time
import matplotlib.pyplot as plt
import numpy as np

from py4mulas.models import Kmodel
from py4mulas.operators import Velocity, OrbitalCurrent
from py4mulas.formulas import KuboFormula
from py4mulas.mu_kernels import KuboKernel

try:
    import kwant
except ImportError as error:
    message = (
        "Kwant is not available. Please install it, "
        "in order to be able to define Kwant Builders. "
    )
    raise ImportError(message) from error


def make_haldane():
    honeycomb = kwant.lattice.honeycomb(1, norbs=1)
    a, b = honeycomb.sublattices
    nnn_hoppings_a = (((-1, 0), a, a), ((0, 1), a, a), ((1, -1), a, a))
    nnn_hoppings_b = (((1, 0), b, b), ((0, -1), b, b), ((-1, 1), b, b))
    nnn_hoppings = nnn_hoppings_a + nnn_hoppings_b

    def onsite(site, M):
        return M if site.family == a else -M

    def nn_hopping(site1, site2, t):
        return t

    def nnn_hopping(site1, site2, t_2):
        return 1j * t_2

    haldane_model = kwant.Builder(kwant.TranslationalSymmetry(*honeycomb.prim_vecs))
    haldane_model[honeycomb.shape(lambda pos: True, (0, 0))] = onsite
    haldane_model[honeycomb.neighbors()] = nn_hopping
    haldane_model[[kwant.builder.HoppingKind(*hopping) for hopping in nnn_hoppings]] = (
        nnn_hopping
    )

    return haldane_model


params = dict(t=-1, t_2=-0.15, M=0)


class OperaKernel:
    def __init__(self, model, alpha, beta, **kwargs):
        self.alpha = alpha
        self.beta = beta
        
        if isinstance(alpha, str):
            self.alpha = Velocity(model, alpha)
        
        if isinstance(beta, str):
            self.beta = Velocity(model, beta)

    def __call__(self, k_args, en, psi, eta):
        alpha = self.alpha(k_args, en, psi)
        beta = self.beta(k_args, en, psi)
        kernel = beta * np.swapaxes(alpha, 1, 2)
        return kernel

def kernel_fun(model, alpha, beta):
    if isinstance(alpha, str):
        alpha = Velocity(model, alpha)
        
    if isinstance(beta, str):
        beta = Velocity(model, beta)

    def fun(k_args, en, psi, eta):
        _alpha = alpha(k_args, en, psi)
        _beta = beta(k_args, en, psi)
        kernel = _beta * np.swapaxes(_alpha, 1, 2)
        return kernel
    def fun2(k_args, en, psi, eta):
        _alpha = alpha(k_args, en, psi)
        _beta = beta(k_args, en, psi)
        kernel = _beta * np.swapaxes(_alpha, 1, 2)
        return kernel
    return [fun, fun2]

def plot_conductivity(nk=100, response="Hall"):
    haldane_model = make_haldane()
    kxs = np.linspace(-np.pi, np.pi, nk)
    kys = kxs
    model = Kmodel(haldane_model, k_1=kxs, k_2=kys, params=params, real_space=False)

    alpha = "x"
    if response == "Hall":
        beta = "y"
    elif response == "OrbitalHall":
        beta = OrbitalCurrent(model, direction="y", gamma="z")

    # We use the prebuilding by setting precomp=True
    # and subdivise the kspace Hamiltonian into momentum chunks
    # the prebuilding facitity does not support variying eta.
    # Once eta is varied the prebuilding restarts.

    options = dict(precomp=True, chunk_size=5000, memmap=True)

    opera_kernel = kernel_fun(model, alpha, beta) #2 * [kernel_fun(model, alpha, beta)] #[OperaKernel(model=model, alpha=alpha, beta=beta)]
    mu_kernel = [KuboKernel("inter_band"), KuboKernel("inter_band")]
    # This should give twice the haldane response

    G = KuboFormula(model, kspace_options=options, mu_kernel=mu_kernel, opera_kernel=opera_kernel)

    energies = np.linspace(-3.5, 3.5, 100)
    t0 = time.time()
    results = []
    for energy in energies:
        results.append(G(energy, temperature=0, eta=0))
    print("taking:", time.time() - t0)
    plt.plot(energies, results)
    plt.show()


def main():
    plot_conductivity(nk=200, response="Hall")


if __name__ == "__main__":
    main()
