import math
import random

SCALE = math.sqrt(0.5)

def prepare(bit, basis):
    if basis == "Z":
        return (1.0, 0j) if bit == 0 else (0j, 1.0)
    if basis == "X":
        return (SCALE, SCALE if bit == 0 else -SCALE)
    raise ValueError("basis must be X or Z")

def measure(state, basis, rng):
    zero = prepare(0, basis)
    overlap = zero[0].conjugate() * state[0] + zero[1].conjugate() * state[1]
    probability_zero = abs(overlap) ** 2
    if probability_zero > 1.0 - 1e-12:
        return 0
    if probability_zero < 1e-12:
        return 1
    return 0 if rng.random() < probability_zero else 1

def qber(pairs):
    if not pairs:
        raise ValueError("QBER is undefined for an empty sifted key")
    return sum(alice != bob for alice, bob in pairs) / len(pairs)

def simulate(alice_bits, alice_bases, bob_bases, attacked, authenticated, seed):
    if not authenticated:
        raise ValueError("BB84 requires an authenticated classical discussion")
    rng = random.Random(seed)
    sifted = []
    evidence = {"eve_wrong_basis": 0, "eve_wrong_basis_errors": 0}
    for alice_bit, alice_basis, bob_basis in zip(alice_bits, alice_bases, bob_bases):
        state = prepare(alice_bit, alice_basis)
        eve_basis = None
        if attacked:
            eve_basis = "X" if rng.randrange(2) else "Z"
            eve_bit = measure(state, eve_basis, rng)
            state = prepare(eve_bit, eve_basis)
        bob_bit = measure(state, bob_basis, rng)
        if alice_basis == bob_basis:
            sifted.append((alice_bit, bob_bit))
            if attacked and eve_basis != alice_basis:
                evidence["eve_wrong_basis"] += 1
                evidence["eve_wrong_basis_errors"] += alice_bit != bob_bit
    return sifted, evidence

schedule = random.Random(84026)
signals = 20000
alice_bits = [schedule.randrange(2) for _ in range(signals)]
alice_bases = ["X" if schedule.randrange(2) else "Z" for _ in range(signals)]
bob_bases = ["X" if schedule.randrange(2) else "Z" for _ in range(signals)]

honest_key, honest_evidence = simulate(alice_bits, alice_bases, bob_bases, False, True, 2601)
attacked_key, attacked_evidence = simulate(alice_bits, alice_bases, bob_bases, True, True, 2602)
honest_qber = qber(honest_key)
attacked_qber = qber(attacked_key)
sift_rate = len(honest_key) / signals
wrong_basis_error_rate = attacked_evidence["eve_wrong_basis_errors"] / attacked_evidence["eve_wrong_basis"]
assert len(honest_key) == len(attacked_key)
assert 0.48 < sift_rate < 0.52
assert honest_qber == 0.0
assert 0.22 < attacked_qber < 0.28
assert 0.46 < wrong_basis_error_rate < 0.54

wrongly_sifted = [(alice, measure(prepare(alice, alice_basis), bob_basis, random.Random(index)),)
                  for index, (alice, alice_basis, bob_basis) in enumerate(zip(alice_bits, alice_bases, bob_bases))
                  if alice_basis != bob_basis]
assert 0.46 < qber(wrongly_sifted) < 0.54

authentication_rejected = False
try:
    simulate([0], ["Z"], ["Z"], False, False, 1)
except ValueError:
    authentication_rejected = True
assert authentication_rejected
empty_sift_rejected = False
try:
    qber([])
except ValueError:
    empty_sift_rejected = True
assert empty_sift_rejected
print(f"PASS: 26 BB84 honest_qber={honest_qber:.3f} attacked_qber={attacked_qber:.3f} sifted={len(honest_key)} wrong_basis_error={wrong_basis_error_rate:.3f}")
