import math

def controlled_x(state, controls, target, qubits):
    if target in controls or len(set(controls)) != len(controls) or any(not 0 <= wire < qubits for wire in (*controls, target)):
        raise ValueError("controls and target must be distinct valid wires")
    control_masks = tuple(1 << (qubits - 1 - wire) for wire in controls)
    target_mask = 1 << (qubits - 1 - target)
    output = [0j] * len(state)
    for basis, amplitude in enumerate(state):
        enabled = all(basis & mask for mask in control_masks)
        output[basis ^ target_mask if enabled else basis] += amplitude
    return output

def basis_state(bits):
    state = [0j] * (1 << len(bits))
    state[int(bits, 2)] = 1.0
    return state

def purity_of_wire(state, wire, qubits):
    mask = 1 << (qubits - 1 - wire)
    rho = [[0j, 0j], [0j, 0j]]
    for left, amplitude_left in enumerate(state):
        for right, amplitude_right in enumerate(state):
            if (left & ~mask) == (right & ~mask):
                bit_left = 1 if left & mask else 0
                bit_right = 1 if right & mask else 0
                rho[bit_left][bit_right] += amplitude_left * amplitude_right.conjugate()
    return sum(abs(value) ** 2 for row in rho for value in row).real

cnot_outputs = set()
for control in (0, 1):
    for target in (0, 1):
        bits = f"{control}{target}"
        output = controlled_x(basis_state(bits), (0,), 1, 2)
        observed = next(index for index, amplitude in enumerate(output) if abs(amplitude) > 0.9)
        expected = 2 * control + (target ^ control)
        assert observed == expected
        cnot_outputs.add(observed)
assert len(cnot_outputs) == 4

toffoli_outputs = set()
for basis in range(8):
    output = controlled_x(basis_state(format(basis, "03b")), (0, 1), 2, 3)
    observed = next(index for index, amplitude in enumerate(output) if abs(amplitude) > 0.9)
    a, b, target = (basis >> 2) & 1, (basis >> 1) & 1, basis & 1
    assert observed == (basis & 0b110) | (target ^ (a & b))
    toffoli_outputs.add(observed)
assert len(toffoli_outputs) == 8

scale = math.sqrt(0.5)
plus_zero = [scale, 0, scale, 0]
bell = controlled_x(plus_zero, (0,), 1, 2)
bell_purity = purity_of_wire(bell, 0, 2)
assert abs(bell_purity - 0.5) < 1e-12

controls_plus_target_zero = [0j] * 8
for index in (0b000, 0b010, 0b100, 0b110):
    controls_plus_target_zero[index] = 0.5
computed = controlled_x(controls_plus_target_zero, (0, 1), 2, 3)
computed_purity = purity_of_wire(computed, 2, 3)
uncomputed = controlled_x(computed, (0, 1), 2, 3)
assert computed_purity < 1.0 and all(abs(a - b) < 1e-12 for a, b in zip(uncomputed, controls_plus_target_zero))
assert abs(purity_of_wire(uncomputed, 2, 3) - 1.0) < 1e-12

invalid_rejected = False
try:
    controlled_x(basis_state("00"), (0,), 0, 2)
except ValueError:
    invalid_rejected = True
assert invalid_rejected
unconditional = controlled_x(plus_zero, (), 1, 2)
assert abs(purity_of_wire(unconditional, 0, 2) - 1.0) < 1e-12 and unconditional != bell
print(f"PASS: 18 controlled-gate tracer verifies {len(cnot_outputs)} CNOT and {len(toffoli_outputs)} Toffoli branches; Bell purity={bell_purity:.3f}, computed-ancilla purity={computed_purity:.3f}, uncompute purity={purity_of_wire(uncomputed, 2, 3):.3f}")
