from qiskit import QuantumCircuit
from qiskit_aer import AerSimulator
import numpy as np
import random


simulator = AerSimulator()
SHOTS = 20000


# =====================================
# CIRCUIT
# =====================================

def create_ghz_circuit(measure_q0=False):

    # 5 qubits, 5 classical bits
    circuit = QuantumCircuit(5, 5)

    # current test circuit
    circuit.x(0)

    circuit.h(1)
    circuit.h(2)
    #circuit.h(3)
    circuit.h(4)
    circuit.cx(1, 0)
    circuit.cx(2, 0)

    # --- NEW: copy q1 → q3 and q2 → q4 ---
    circuit.cx(1, 3)
    circuit.cx(2, 3)
    circuit.cx(0, 3)

    circuit.barrier()

    return circuit


# =====================================
# CHSH BASIS
# =====================================

def measure_in_basis(circuit, qubit, angle):
    circuit.ry(-2 * angle, qubit)


# =====================================
# SPLIT q0
# =====================================

def split_counts_by_q0(counts):
    q0_0 = {}
    q0_1 = {}

    for key, value in counts.items():
        # qiskit order: c4 c3 c2 c1 c0
        q2q1 = key[2:4]          # c2 c1
        q0 = key[4]               # c0

        if q0 == "0":
            q0_0[q2q1] = q0_0.get(q2q1, 0) + value
        else:
            q0_1[q2q1] = q0_1.get(q2q1, 0) + value

    return q0_0, q0_1


# =====================================
# CORRELATION
# =====================================

def calculate_correlation(counts):
    total = sum(counts.values())
    if total == 0:
        return 0

    corr = 0
    for key, value in counts.items():
        q2 = key[0]
        q1 = key[1]
        if q1 == q2:
            corr += value
        else:
            corr -= value

    return corr / total


# =====================================
# TABLE
# =====================================

def print_table(rows):
    print("\n")
    print("=" * 110)
    print(
        f"{'TEST':8}"
        f"{'q0':5}"
        f"{'00':7}"
        f"{'01':7}"
        f"{'10':7}"
        f"{'11':7}"
        f"{'N':8}"
        f"{'E(q1,q2)':12}"
    )
    print("-" * 110)

    for r in rows:
        print(
            f"{r['test']:8}"
            f"{r['q0']:5}"
            f"{r['00']:7}"
            f"{r['01']:7}"
            f"{r['10']:7}"
            f"{r['11']:7}"
            f"{r['N']:8}"
            f"{r['corr']:12.6f}"
        )
    print("=" * 110)


# =====================================
# CHSH SETTINGS
# =====================================

bases = [
    ("A0B0", 0, np.pi / 8),
    ("A0B1", 0, -np.pi / 8),
    ("A1B0", np.pi / 4, np.pi / 8),
    ("A1B1", np.pi / 4, -np.pi / 8)
]


# =====================================
# RUN
# =====================================

def run_chsh(measure_q0, rnd):

    results_all = []
    results_q0_0 = []
    results_q0_1 = []
    table = []

    print("\n")
    print("=" * 40)
    if measure_q0:
        print("q0 MEASURED")
    else:
        print("q0 NOT MEASURED")
    print("=" * 40)

    for name, a, b in bases:

        circuit = create_ghz_circuit(measure_q0)

        # measurement bases on q1 and q2
        measure_in_basis(circuit, 1, a)
        measure_in_basis(circuit, 2, b)

        circuit.barrier()

        # measure q1 and q2
        circuit.measure(1, 1)
        circuit.barrier()
        circuit.measure(2, 2)
        circuit.barrier()

        # measure q0 if requested
        meas = measure_q0
        if meas:
            circuit.measure(0, 0)

        # measure the new ancillas q3 and q4 (optional, for checking)
        circuit.measure(3, 3)
        circuit.measure(4, 4)

        print("\n", name)
        print(circuit.draw())

        result = simulator.run(circuit, shots=SHOTS).result()
        counts = result.get_counts()

        # =====================================
        # CHECK q0=1 AND q3=1
        # =====================================

        q0_q3_11 = 0

        total_shots = 0
        for key, value in counts.items():
            q0 = key[4]
            q3 = key[1]

            if q0 == "1":
                total_shots += value
                if q3 == "1":
                    q0_q3_11 += value

        if total_shots > 0:
            print(
                f"q0=1 -> q3=1: {q0_q3_11} / {total_shots}"
                f"  ({100*q0_q3_11/total_shots:.4f}%)"
            )
        else:
            print("\nq0 ALWAYS 0")

        print("\nCOUNTS")
        print(counts)

        # --- total correlation on q1,q2 ---
        clean = {}
        for k, v in counts.items():
            # take only bits of q2 and q1 (positions 2 and 1 in the 5-bit string)
            q2q1 = k[2:4]
            clean[q2q1] = clean.get(q2q1, 0) + v

        corr = calculate_correlation(clean)
        results_all.append(corr)

        if not meas:
            table.append({
                "test": name,
                "q0": "--",
                "00": clean.get("00", 0),
                "01": clean.get("01", 0),
                "10": clean.get("10", 0),
                "11": clean.get("11", 0),
                "N": sum(clean.values()),
                "corr": corr
            })

        if meas:
            c0, c1 = split_counts_by_q0(counts)

            corr0 = calculate_correlation(c0)
            corr1 = calculate_correlation(c1)

            results_q0_0.append(corr0)
            results_q0_1.append(corr1)

            for label, c in [("0", c0), ("1", c1)]:
                table.append({
                    "test": name,
                    "q0": label,
                    "00": c.get("00", 0),
                    "01": c.get("01", 0),
                    "10": c.get("10", 0),
                    "11": c.get("11", 0),
                    "N": sum(c.values()),
                    "corr": calculate_correlation(c)
                })

    print_table(table)

    # Total CHSH
    S = results_all[0] + results_all[1] + results_all[2] - results_all[3]
    print("\nTotal CHSH =", S)

    if measure_q0:
        S0 = results_q0_0[0] + results_q0_0[1] + results_q0_0[2] - results_q0_0[3]
        S1 = results_q0_1[0] + results_q0_1[1] + results_q0_1[2] - results_q0_1[3]

        print("\nConditional CHSH")
        print("q0=0 :", S0)
        print("q0=1 :", S1)

        S0_alt = -results_q0_0[0] - results_q0_0[1] + results_q0_0[2] - results_q0_0[3]
        print("\nAlternative CHSH q0=0")
        print("S0_alt =", S0_alt)
        print("2*sqrt(2) =", 2 * np.sqrt(2))


# =====================================
# MAIN
# =====================================

run_chsh(False, True)
run_chsh(True, True)

run_chsh(False, False)
run_chsh(True, False)
