import math

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

def pauli(state, qubit, name):
    mask = 1 << (1 - qubit)
    output = [0j] * 4
    for basis, amplitude in enumerate(state):
        if name == "X":
            output[basis ^ mask] += amplitude
        elif name == "Z":
            output[basis] += -amplitude if basis & mask else amplitude
        else:
            raise ValueError("only Pauli X and Z are supported")
    return output

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

def prepare(shared_entanglement):
    state = [1.0, 0j, 0j, 0j]
    if shared_entanglement:
        state = cnot(h(state, 0), 0, 1)
    return state

def encode(state, message):
    first, second = message
    if second:
        state = pauli(state, 0, "X")
    if first:
        state = pauli(state, 0, "Z")
    return state

def decode(state, reversed_order=False):
    state = h(state, 0) if reversed_order else cnot(state, 0, 1)
    state = cnot(state, 0, 1) if reversed_order else h(state, 0)
    return state

def distribution(state):
    return {format(index, "02b"): round(abs(amplitude) ** 2, 12)
            for index, amplitude in enumerate(state) if abs(amplitude) > 1e-12}

def inner(left, right):
    return sum(a.conjugate() * b for a, b in zip(left, right))

messages = [(0, 0), (0, 1), (1, 0), (1, 1)]
bell = prepare(True)
assert all(abs(a - b) < 1e-12 for a, b in zip(bell, [math.sqrt(0.5), 0, 0, math.sqrt(0.5)]))
codewords = {message: encode(bell, message) for message in messages}
for left in messages:
    for right in messages:
        overlap = abs(inner(codewords[left], codewords[right]))
        assert abs(overlap - (1.0 if left == right else 0.0)) < 1e-12

traces = {}
for message in messages:
    encoded = codewords[message]
    decoded = decode(encoded)
    observed = distribution(decoded)
    expected = {"".join(str(bit) for bit in message): 1.0}
    assert observed == expected
    traces[message] = (bell, encoded, decoded, observed)
assert len({next(iter(trace[3])) for trace in traces.values()}) == 4

without_ebit = {message: distribution(decode(encode(prepare(False), message))) for message in messages}
assert without_ebit[(0, 0)] == without_ebit[(1, 0)]
assert without_ebit[(0, 1)] == without_ebit[(1, 1)]
assert len({tuple(sorted(value.items())) for value in without_ebit.values()}) == 2

wrong_decode = {message: distribution(decode(codewords[message], True)) for message in messages}
assert any(wrong_decode[message] != {"".join(str(bit) for bit in message): 1.0} for message in messages)
resources = {"transmitted_qubits": 1, "shared_ebits_consumed": 1, "decoded_classical_bits": 2}
assert resources == {"transmitted_qubits": 1, "shared_ebits_consumed": 1, "decoded_classical_bits": 2}
print(f"PASS: 24 superdense circuit codewords={len(codewords)} decoded={len(traces)} no_ebit_outputs={len({tuple(sorted(value.items())) for value in without_ebit.values()})}")
