import math
import random

def normalize(alpha, beta):
    norm = math.sqrt(abs(alpha) ** 2 + abs(beta) ** 2)
    if norm < 1e-15:
        raise ValueError("the zero vector is not a quantum state")
    return alpha / norm, beta / norm

def h(state, qubit):
    mask = 1 << (2 - qubit)
    output = state[:]
    scale = math.sqrt(0.5)
    for basis in range(8):
        if basis & mask == 0:
            paired = basis | mask
            output[basis] = (state[basis] + state[paired]) * scale
            output[paired] = (state[basis] - state[paired]) * scale
    return output

def cnot(state, control, target):
    control_mask = 1 << (2 - control)
    target_mask = 1 << (2 - target)
    output = [0j] * 8
    for basis, amplitude in enumerate(state):
        output[basis ^ target_mask if basis & control_mask else basis] += amplitude
    return output

def alice_circuit(alpha, beta):
    state = [0j] * 8
    state[0] = alpha
    state[4] = beta
    state = cnot(h(state, 1), 1, 2)
    state = h(cnot(state, 0, 1), 0)
    return state

def receiver_branch(state, first, second):
    amplitudes = [state[(first << 2) | (second << 1) | bob] for bob in (0, 1)]
    probability = sum(abs(value) ** 2 for value in amplitudes)
    if probability < 1e-15:
        raise ValueError("impossible measurement branch")
    return probability, tuple(value / math.sqrt(probability) for value in amplitudes)

def correct(receiver, first, second):
    zero, one = receiver
    if second:
        zero, one = one, zero
    if first:
        one = -one
    return zero, one

def fidelity(left, right):
    overlap = left[0].conjugate() * right[0] + left[1].conjugate() * right[1]
    return abs(overlap) ** 2

def reduced_receiver(branches):
    rho = [[0j, 0j], [0j, 0j]]
    for probability, receiver in branches:
        for row in (0, 1):
            for column in (0, 1):
                rho[row][column] += probability * receiver[row] * receiver[column].conjugate()
    return rho

rng = random.Random(2501)
states = [normalize(1, 0), normalize(0, 1), normalize(1, 1j), normalize(3 - 2j, -1 + 4j)]
for _ in range(12):
    states.append(normalize(complex(rng.uniform(-1, 1), rng.uniform(-1, 1)),
                            complex(rng.uniform(-1, 1), rng.uniform(-1, 1))))

minimum_fidelity = 1.0
for psi in states:
    state = alice_circuit(*psi)
    branches = []
    for first in (0, 1):
        for second in (0, 1):
            probability, receiver = receiver_branch(state, first, second)
            assert abs(probability - 0.25) < 1e-12
            corrected = correct(receiver, first, second)
            branch_fidelity = fidelity(psi, corrected)
            minimum_fidelity = min(minimum_fidelity, branch_fidelity)
            assert abs(branch_fidelity - 1.0) < 1e-12
            branches.append((probability, receiver))
    rho = reduced_receiver(branches)
    assert abs(rho[0][0] - 0.5) < 1e-12 and abs(rho[1][1] - 0.5) < 1e-12
    assert abs(rho[0][1]) < 1e-12 and abs(rho[1][0]) < 1e-12
    unavailable_fidelity = sum((psi[row].conjugate() * rho[row][column] * psi[column]).real
                               for row in (0, 1) for column in (0, 1))
    assert abs(unavailable_fidelity - 0.5) < 1e-12

witness = normalize(1, 2j)
witness_state = alice_circuit(*witness)
mutated_failures = 0
for first in (0, 1):
    for second in (0, 1):
        _, receiver = receiver_branch(witness_state, first, second)
        wrong = correct(receiver, second, first)
        mutated_failures += fidelity(witness, wrong) < 1.0 - 1e-9
assert mutated_failures > 0
rejected_zero = False
try:
    normalize(0, 0)
except ValueError:
    rejected_zero = True
assert rejected_zero
print(f"PASS: 25 teleportation states={len(states)} minimum_fidelity={minimum_fidelity:.12f} swapped_correction_failures={mutated_failures}")
