import math

def h(state, qubit, count):
    mask = 1 << (count - 1 - qubit)
    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 query_oracle(state, function, input_qubits):
    output = [0j] * len(state)
    for basis, amplitude in enumerate(state):
        x = basis >> 1
        y = basis & 1
        destination = (x << 1) | (y ^ function(x))
        output[destination] += amplitude
    return output

def one_query_state(function, input_qubits, target_one=True, query_count=1):
    total = input_qubits + 1
    state = [0j] * (1 << total)
    state[1 if target_one else 0] = 1.0
    for qubit in range(total):
        state = h(state, qubit, total)
    before_oracle = state[:]
    for _ in range(query_count):
        state = query_oracle(state, function, input_qubits)
    after_oracle = state[:]
    for qubit in range(input_qubits):
        state = h(state, qubit, total)
    probabilities = [0.0] * (1 << input_qubits)
    for basis, amplitude in enumerate(state):
        probabilities[basis >> 1] += abs(amplitude) ** 2
    return before_oracle, after_oracle, probabilities

def assert_kickback(after_oracle, function, input_qubits):
    expected_size = math.sqrt(1 / (1 << (input_qubits + 1)))
    for x in range(1 << input_qubits):
        phase = -1 if function(x) else 1
        assert abs(after_oracle[(x << 1)] - phase * expected_size) < 1e-12
        assert abs(after_oracle[(x << 1) | 1] + phase * expected_size) < 1e-12

def classify_dj(table):
    size = len(table)
    input_qubits = size.bit_length() - 1
    if 1 << input_qubits != size:
        raise ValueError("truth table length must be a power of two")
    ones = sum(table)
    if ones not in (0, size, size // 2):
        raise ValueError("Deutsch-Jozsa promise violated")
    function = lambda x: table[x]
    _, kicked, probabilities = one_query_state(function, input_qubits)
    assert_kickback(kicked, function, input_qubits)
    measured_zero = probabilities[0] > 1.0 - 1e-12
    return "constant" if measured_zero else "balanced", probabilities

deutsch_tables = 0
for encoded in range(4):
    table = tuple((encoded >> x) & 1 for x in range(2))
    classification, probabilities = classify_dj(table)
    assert classification == ("constant" if sum(table) in (0, 2) else "balanced")
    deutsch_tables += 1
assert deutsch_tables == 4

promised_tables = 0
for encoded in range(1 << 8):
    table = tuple((encoded >> x) & 1 for x in range(8))
    if sum(table) in (0, 4, 8):
        classification, probabilities = classify_dj(table)
        assert classification == ("constant" if sum(table) in (0, 8) else "balanced")
        if classification == "balanced":
            assert probabilities[0] < 1e-12
        promised_tables += 1
assert promised_tables == 72

promise_rejected = False
try:
    classify_dj((1, 0, 0, 0, 0, 0, 0, 0))
except ValueError:
    promise_rejected = True
assert promise_rejected

bv_cases = 0
for input_qubits in range(1, 6):
    for secret in range(1 << input_qubits):
        for offset in (0, 1):
            function = lambda x, secret=secret, offset=offset: ((x & secret).bit_count() & 1) ^ offset
            _, kicked, probabilities = one_query_state(function, input_qubits)
            assert_kickback(kicked, function, input_qubits)
            assert abs(probabilities[secret] - 1.0) < 1e-12
            assert sum(value > 1e-12 for value in probabilities) == 1
            ledger = {
                "oracle_queries": 1,
                "hadamards": 2 * input_qubits + 1,
                "compiled_oracle_cnot": secret.bit_count(),
                "compiled_oracle_x": offset,
            }
            assert ledger["oracle_queries"] == 1
            assert ledger["hadamards"] == 2 * input_qubits + 1
            assert ledger["compiled_oracle_cnot"] + ledger["compiled_oracle_x"] == secret.bit_count() + offset
            bv_cases += 1

witness_secret = 0b1011
witness = lambda x: (x & witness_secret).bit_count() & 1
_, _, correct = one_query_state(witness, 4)
_, _, target_mutation = one_query_state(witness, 4, target_one=False)
_, _, double_query_mutation = one_query_state(witness, 4, query_count=2)
assert correct[witness_secret] > 1.0 - 1e-12
assert target_mutation[0] > 1.0 - 1e-12 and target_mutation[witness_secret] < 1e-12
assert double_query_mutation[0] > 1.0 - 1e-12 and double_query_mutation[witness_secret] < 1e-12
print(f"PASS: 28 one-query algorithms Deutsch={deutsch_tables} DJ={promised_tables} BV={bv_cases} witness_success={correct[witness_secret]:.1f} target_mutant_success={target_mutation[witness_secret]:.1f}")
