import math

def validate_layout(qubits, layout):
    if tuple(sorted(layout)) != tuple(range(qubits)):
        raise ValueError("layout must map each logical wire exactly once")

def wire_mask(wire, qubits, layout):
    validate_layout(qubits, layout)
    if not 0 <= wire < qubits:
        raise ValueError("wire out of range")
    return 1 << (qubits - 1 - layout[wire])

def encode_bits(bits, layout):
    qubits = len(bits)
    index = 0
    for wire, bit in enumerate(bits):
        if bit:
            index |= wire_mask(wire, qubits, layout)
    return index

def decode_bits(index, qubits, layout):
    return tuple(1 if index & wire_mask(wire, qubits, layout) else 0 for wire in range(qubits))

def apply_x(state, wire, qubits, layout):
    mask = wire_mask(wire, qubits, layout)
    output = [0j] * len(state)
    for basis, amplitude in enumerate(state):
        output[basis ^ mask] += amplitude
    return output

def apply_h(state, wire, qubits, layout):
    mask = wire_mask(wire, qubits, layout)
    output = state[:]
    scale = math.sqrt(0.5)
    for basis in range(len(state)):
        if basis & mask == 0:
            paired = basis | mask
            output[basis] = (state[basis] + state[paired]) * scale
            output[paired] = (state[basis] - state[paired]) * scale
    return output

def apply_cnot(state, control, target, qubits, layout):
    if control == target:
        raise ValueError("control and target must differ")
    control_mask = wire_mask(control, qubits, layout)
    target_mask = wire_mask(target, qubits, layout)
    output = [0j] * len(state)
    for basis, amplitude in enumerate(state):
        output[basis ^ target_mask if basis & control_mask else basis] += amplitude
    return output

layouts = ((0, 1, 2), (2, 1, 0), (1, 2, 0))
local_gate_cases = 0
for layout in layouts:
    for wire in range(3):
        state = [0j] * 8
        state[encode_bits((0, 0, 0), layout)] = 1.0
        output = apply_x(state, wire, 3, layout)
        observed_index = next(index for index, amplitude in enumerate(output) if abs(amplitude) > 0.9)
        expected = tuple(1 if index == wire else 0 for index in range(3))
        assert decode_bits(observed_index, 3, layout) == expected
        local_gate_cases += 1

bell_supports = []
for layout in layouts:
    state = [0j] * 8
    state[encode_bits((0, 0, 0), layout)] = 1.0
    state = apply_cnot(apply_h(state, 0, 3, layout), 0, 2, 3, layout)
    support = {decode_bits(index, 3, layout): round(abs(amplitude) ** 2, 12)
               for index, amplitude in enumerate(state) if abs(amplitude) > 1e-12}
    assert support == {(0, 0, 0): 0.5, (1, 0, 1): 0.5}
    bell_supports.append(support)

for layout in ((0, 1, 2, 3), (3, 2, 1, 0)):
    state = [0j] * 16
    state[encode_bits((1, 0, 0, 0), layout)] = 1.0
    output = apply_cnot(state, 0, 3, 4, layout)
    observed = next(index for index, amplitude in enumerate(output) if abs(amplitude) > 0.9)
    assert decode_bits(observed, 4, layout) == (1, 0, 0, 1)

ket_01_msb = encode_bits((0, 1), (0, 1))
ket_01_lsb = encode_bits((0, 1), (1, 0))
assert (ket_01_msb, ket_01_lsb) == (1, 2)
bad_layout_rejected = False
try:
    encode_bits((0, 1, 0), (0, 0, 2))
except ValueError:
    bad_layout_rejected = True
assert bad_layout_rejected
print(f"PASS: 19 wire-scope suite checks {local_gate_cases} local gates, {len(bell_supports)} layout-stable Bell traces, and 2 nonadjacent CNOTs; |01> indices={ket_01_msb}/{ket_01_lsb}")
