import math
import random

def dot2(left, right):
    return (left & right).bit_count() & 1

def promised_oracle(secret):
    return lambda x: min(x, x ^ secret)

def promise_holds(function, secret, input_qubits):
    outputs = {}
    for x in range(1 << input_qubits):
        value = function(x)
        if value in outputs and outputs[value] != (x ^ secret):
            return False
        outputs[value] = x
        if function(x ^ secret) != value:
            return False
    return len(outputs) == 1 << (input_qubits - 1)

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 simon_distribution(function, input_qubits):
    total = 2 * input_qubits
    state = [0j] * (1 << total)
    state[0] = 1.0
    for qubit in range(input_qubits):
        state = h(state, qubit, total)
    queried = [0j] * len(state)
    output_mask = (1 << input_qubits) - 1
    for basis, amplitude in enumerate(state):
        x = basis >> input_qubits
        output = basis & output_mask
        queried[(x << input_qubits) | (output ^ function(x))] += amplitude
    state = queried
    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 >> input_qubits] += abs(amplitude) ** 2
    return probabilities

def rref(rows, input_qubits):
    reduced = [row for row in rows if row]
    pivots = []
    pivot_row = 0
    for column in range(input_qubits):
        bit = 1 << (input_qubits - 1 - column)
        found = next((index for index in range(pivot_row, len(reduced)) if reduced[index] & bit), None)
        if found is None:
            continue
        reduced[pivot_row], reduced[found] = reduced[found], reduced[pivot_row]
        for index in range(len(reduced)):
            if index != pivot_row and reduced[index] & bit:
                reduced[index] ^= reduced[pivot_row]
        pivots.append(column)
        pivot_row += 1
        if pivot_row == len(reduced):
            break
    return reduced[:pivot_row], pivots

def recover(rows, input_qubits):
    reduced, pivots = rref(rows, input_qubits)
    if len(pivots) != input_qubits - 1:
        raise ValueError("collect samples until GF(2) rank is n-1")
    free_columns = [column for column in range(input_qubits) if column not in pivots]
    if len(free_columns) != 1:
        raise ValueError("expected a one-dimensional nullspace")
    candidate = 1 << (input_qubits - 1 - free_columns[0])
    for row, column in reversed(list(zip(reduced, pivots))):
        pivot_bit = 1 << (input_qubits - 1 - column)
        if dot2(row ^ pivot_bit, candidate):
            candidate |= pivot_bit
    if candidate == 0 or any(dot2(row, candidate) for row in rows):
        raise ValueError("invalid nullspace recovery")
    return candidate

def validate_samples(rows, secret):
    if any(dot2(row, secret) for row in rows):
        raise ValueError("sample violates y dot s = 0")

tested = 0
largest_sample_log = 0
for input_qubits in range(1, 6):
    for secret in range(1, 1 << input_qubits):
        function = promised_oracle(secret)
        assert promise_holds(function, secret, input_qubits)
        probabilities = simon_distribution(function, input_qubits)
        support = [y for y, probability in enumerate(probabilities) if probability > 1e-12]
        expected_support = [y for y in range(1 << input_qubits) if dot2(y, secret) == 0]
        assert support == expected_support
        expected_probability = 1 / (1 << (input_qubits - 1))
        assert all(abs(probabilities[y] - expected_probability) < 1e-12 for y in support)
        assert all(probabilities[y] < 1e-12 for y in range(len(probabilities)) if y not in support)

        rng = random.Random(29000 + 101 * input_qubits + secret)
        rows = []
        while len(rref(rows, input_qubits)[1]) < input_qubits - 1:
            rows.append(rng.choice(support))
            if len(rows) > 200:
                raise AssertionError("seeded sampling failed to reach rank")
        validate_samples(rows, secret)
        recovered = recover(rows, input_qubits)
        assert recovered == secret
        assert all(function(x) == function(x ^ recovered) for x in range(1 << input_qubits))
        largest_sample_log = max(largest_sample_log, len(rows))
        tested += 1

insufficient_rejected = False
try:
    recover([0b0100], 4)
except ValueError:
    insufficient_rejected = True
assert insufficient_rejected

invalid_sample_rejected = False
try:
    validate_samples([0b0001], 0b1011)
except ValueError:
    invalid_sample_rejected = True
assert invalid_sample_rejected

good = promised_oracle(0b1011)
bad = lambda x: 0b1111 if x == 0 else good(x)
assert promise_holds(good, 0b1011, 4)
assert not promise_holds(bad, 0b1011, 4)
bad_probabilities = simon_distribution(bad, 4)
assert any(bad_probabilities[y] > 1e-12 and dot2(y, 0b1011) for y in range(16))
print(f"PASS: 29 Simon recovered={tested} max_samples={largest_sample_log} bad_forbidden_mass={sum(bad_probabilities[y] for y in range(16) if dot2(y, 0b1011)):.6f}")
