#!/usr/bin/env python3
"""
Libreria comune degli script che generano i dataset del quaderno n. 3.

Dipendenze: Python 3 e numpy. Nessun SDK quantistico: ogni numero del volume
si rigenera con `make dati` su una macchina qualunque, senza accesso a
hardware IBM. Gli esperimenti su hardware sono descritti nel volume e
citati dalle fonti; i dati qui prodotti sono simulazioni esatte su sistemi
piccoli, costruite per rendere visibili i meccanismi.

Convenzioni (v. capitolo «Notazione e convenzioni»):
    - qubit indicizzati da 0; lo stato iniziale è |0...0>;
    - una stringa di Pauli su n qubit è una coppia di interi (x, z) nella
      rappresentazione simplettica: il qubit j porta X se è acceso il solo
      bit j di x, Z se è acceso il solo bit j di z, Y se sono accesi entrambi;
    - le rotazioni sono R_P(theta) = exp(-i theta P / 2);
    - il circuito di riferimento è l'Ising a calcio (kicked Ising): uno
      strato di RZZ(theta_J) sui lati del grafo seguito da RX(theta_h) su
      ogni qubit, ripetuto L volte.
"""

import math
import os

import numpy as np

HERE = os.path.dirname(os.path.abspath(__file__))


def percorso(nome):
    return os.path.join(HERE, nome)


def scrivi_dat(nome, intestazione, righe, fmt="{:.9g}"):
    """Scrive un .dat per pgfplots: prima riga di commento con le colonne."""
    with open(percorso(nome), "w") as f:
        f.write("# " + intestazione + "\n")
        for r in righe:
            f.write(" ".join(fmt.format(v) if not isinstance(v, str) else v
                             for v in r) + "\n")


def scrivi_txt(nome, intestazione, righe, nota=None):
    """Scrive una tabella già formattata all'italiana, celle separate da ' & '."""
    with open(percorso(nome), "w") as f:
        f.write(" & ".join(intestazione) + "\n")
        for r in righe:
            f.write(" & ".join(r) + "\n")
        if nota:
            f.write("% " + nota + "\n")


def it(x, cifre=2):
    """Numero in notazione italiana per LaTeX: virgola decimale, punto migliaia."""
    s = f"{x:,.{cifre}f}"
    s = s.replace(",", "X").replace(".", "{,}").replace("X", ".")
    return s


def it_sci(x, cifre=2):
    """Notazione scientifica italiana per LaTeX: 1{,}23\\cdot10^{4}."""
    if x == 0:
        return "0"
    e = int(math.floor(math.log10(abs(x))))
    m = x / 10 ** e
    if round(m, cifre) >= 10:
        m /= 10
        e += 1
    return f"{it(m, cifre)}\\cdot10^{{{e}}}"


# ============================================================ grafi
def catena(n, chiusa=False):
    lati = [(i, i + 1) for i in range(n - 1)]
    if chiusa:
        lati.append((n - 1, 0))
    return lati


def colorazione_lati(lati):
    """Partiziona i lati in strati senza qubit in comune (colorazione greedy)."""
    strati = []
    for e in lati:
        for s in strati:
            if all(e[0] not in f and e[1] not in f for f in s):
                s.append(e)
                break
        else:
            strati.append([e])
    return strati


# ============================================================ vettore di stato
def stato_zero(n):
    psi = np.zeros(2 ** n, dtype=complex)
    psi[0] = 1.0
    return psi


def _bit(n, q):
    """Vettore 0/1 del bit del qubit q per ogni indice di base (qubit 0 = bit meno significativo)."""
    idx = np.arange(2 ** n)
    return (idx >> q) & 1


def applica_rzz(psi, n, a, b, theta):
    za = 1 - 2 * _bit(n, a)
    zb = 1 - 2 * _bit(n, b)
    fase = np.exp(-1j * theta / 2 * za * zb)
    return psi * fase


def applica_rx(psi, n, q, theta):
    c, s = math.cos(theta / 2), math.sin(theta / 2)
    psi = psi.reshape([2] * n)
    ax = n - 1 - q  # asse numpy del qubit q (ordine big-endian degli assi)
    psi = np.moveaxis(psi, ax, 0)
    a0, a1 = psi[0].copy(), psi[1].copy()
    psi[0] = c * a0 - 1j * s * a1
    psi[1] = -1j * s * a0 + c * a1
    psi = np.moveaxis(psi, 0, ax)
    return psi.reshape(-1)


def kicked_ising_esatto(n, lati, theta_j, theta_h, passi, osservabile_q):
    """<Z_q>(t) esatto per t = 0..passi sul circuito di Ising a calcio."""
    psi = stato_zero(n)
    zq = 1 - 2 * _bit(n, osservabile_q)
    out = [float(np.real(np.vdot(psi, zq * psi)))]
    for _ in range(passi):
        for (a, b) in lati:
            psi = applica_rzz(psi, n, a, b, theta_j)
        for q in range(n):
            psi = applica_rx(psi, n, q, theta_h)
        out.append(float(np.real(np.vdot(psi, zq * psi))))
    return out


# ============================================================ algebra di Pauli
def peso(x, z):
    return bin(x | z).count("1")


def prodotto(p1, p2):
    """Prodotto di due Pauli (x1,z1)*(x2,z2) = i^k (x1^x2, z1^z2): restituisce (x, z, k mod 4).

    Convenzione: la stringa (x,z) rappresenta il prodotto tensore di
    P_j = i^{x_j z_j} X^{x_j} Z^{z_j}, cioè Y = iXZ. Il calcolo della fase
    segue Aaronson–Gottesman, qubit per qubit."""
    x1, z1 = p1
    x2, z2 = p2
    k = 0
    xs, zs = x1 | x2 | z1 | z2, 0
    j = 0
    m = x1 | z1 | x2 | z2
    while m:
        if m & 1:
            a1, b1 = (x1 >> j) & 1, (z1 >> j) & 1
            a2, b2 = (x2 >> j) & 1, (z2 >> j) & 1
            k += _g(a1, b1, a2, b2)
        m >>= 1
        j += 1
    return x1 ^ x2, z1 ^ z2, k % 4


def _g(x1, z1, x2, z2):
    """Esponente di i nel prodotto di due Pauli a un qubit (Aaronson–Gottesman)."""
    if x1 == 0 and z1 == 0:
        return 0
    if x1 == 1 and z1 == 1:
        return z2 - x2
    if x1 == 1 and z1 == 0:
        return z2 * (2 * x2 - 1)
    return x2 * (1 - 2 * z2)


def anticommutano(p1, p2):
    x1, z1 = p1
    x2, z2 = p2
    return (bin(x1 & z2).count("1") + bin(z1 & x2).count("1")) % 2 == 1


def coniuga_rotazione(operatore, generatore, theta, soglia=0.0, peso_max=None):
    """Immagine di Heisenberg U^† O U con U = exp(-i theta G/2), O = somma di Pauli.

    `operatore` è un dict {(x,z): coefficiente reale}. Se P anticommuta con G:
        P -> cos(theta) P + sin(theta) (i G P),
    dove i G P è una Pauli hermitiana con segno ±1. I termini con |c| sotto
    `soglia` o con peso oltre `peso_max` vengono scartati (troncamento).
    Restituisce (nuovo operatore, norma 1 dei coefficienti scartati)."""
    nuovo = {}
    scartato = 0.0
    c_t, s_t = math.cos(theta), math.sin(theta)
    for p, c in operatore.items():
        if not anticommutano(p, generatore):
            nuovo[p] = nuovo.get(p, 0.0) + c
            continue
        nuovo[p] = nuovo.get(p, 0.0) + c * c_t
        x, z, k = prodotto(generatore, p)   # G P = i^k Q
        # i G P = i^{k+1} Q; hermitiano => k+1 pari => segno = (-1)^{(k+1)/2}
        segno = 1.0 if (k + 1) % 4 == 0 else -1.0
        q = (x, z)
        nuovo[q] = nuovo.get(q, 0.0) + c * s_t * segno
    out = {}
    for p, c in nuovo.items():
        if abs(c) <= soglia or (peso_max is not None and peso(*p) > peso_max):
            scartato += abs(c)
            continue
        out[p] = c
    return out, scartato


def propaga_pauli(n, lati, theta_j, theta_h, passi, osservabile_q,
                  soglia=0.0, peso_max=None, smorzamento=1.0):
    """Propagazione di Pauli (immagine di Heisenberg) di Z_q nel circuito di Ising a calcio.

    Il circuito di t passi è U = (K V)^t con V = prod RZZ, K = prod RX.
    <Z_q(t)> = <0| U^† Z U |0>. Si coniuga strato per strato partendo
    dall'ultimo applicato. Restituisce, per t = 0..passi, il valore stimato,
    il numero di termini e la norma scartata cumulata."""
    risultati = []
    for t in range(passi + 1):
        op = {(0, 1 << osservabile_q): 1.0}
        scart = 0.0
        massimo = 1
        for _ in range(t):
            if smorzamento != 1.0:
                op = {p: c * smorzamento ** peso(*p) for p, c in op.items()}
            # ultimo strato applicato: RX su tutti i qubit
            for q in range(n):
                op, s = coniuga_rotazione(op, (1 << q, 0), theta_h, soglia, peso_max)
                scart += s
            for (a, b) in lati:
                op, s = coniuga_rotazione(op, (0, (1 << a) | (1 << b)), theta_j,
                                          soglia, peso_max)
                scart += s
            massimo = max(massimo, len(op))
        val = sum(c for (x, z), c in op.items() if x == 0)
        risultati.append((t, val, len(op), massimo, scart))
    return risultati


# ============================================================ statistica
def shot_per_precisione(eps, varianza_singolo=1.0, fattore=1.0):
    """Numero di ripetizioni per errore standard eps, con amplificazione `fattore` della varianza."""
    return fattore * varianza_singolo / eps ** 2
