import math

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

def matvec(matrix, vector):
    return tuple(sum(value * vector[column] for column, value in enumerate(row)) for row in matrix)

def matmul(left, right):
    return tuple(tuple(sum(left[row][k] * right[k][column] for k in range(len(right)))
                       for column in range(len(right[0]))) for row in range(len(left)))

def projector(ket):
    norm = inner(ket, ket).real
    if abs(norm - 1.0) > 1e-12:
        raise ValueError("projector ket must be normalized")
    return tuple(tuple(ket[row] * ket[column].conjugate() for column in range(len(ket)))
                 for row in range(len(ket)))

def condition(matrix, state):
    projected = matvec(matrix, state)
    probability = inner(projected, projected).real
    if probability < 1e-15:
        raise ValueError("cannot normalize a zero-probability branch")
    return probability, tuple(value / math.sqrt(probability) for value in projected)

scale = math.sqrt(0.5)
kets = ((1 + 0j, 0j), (scale, 1j * scale))
states = ((0.6, 0.8), (scale, scale), (scale, -1j * scale))
max_idempotence_error = 0.0
probabilities = []
for ket in kets:
    matrix = projector(ket)
    squared = matmul(matrix, matrix)
    max_idempotence_error = max(max_idempotence_error,
                                *(abs(squared[row][column] - matrix[row][column])
                                  for row in range(2) for column in range(2)))
    for state in states:
        probability = inner(state, matvec(matrix, state)).real
        assert -1e-12 <= probability <= 1.0 + 1e-12
        probabilities.append(probability)
        if probability > 1e-12:
            branch_probability, conditional = condition(matrix, state)
            assert abs(branch_probability - probability) < 1e-12
            assert abs(inner(conditional, conditional) - 1.0) < 1e-12
assert max_idempotence_error < 1e-12

for left in states:
    for right in states:
        assert abs(inner(left, right) - inner(right, left).conjugate()) < 1e-12
assert abs(inner(kets[0], kets[1])) == scale

zero_branch_rejected = False
try:
    condition(projector((1, 0)), (0, 1))
except ValueError:
    zero_branch_rejected = True
assert zero_branch_rejected
bad_projector_rejected = False
try:
    projector((2, 0))
except ValueError:
    bad_projector_rejected = True
assert bad_projector_rejected
naive = lambda left, right: sum(a * b for a, b in zip(left, right))
assert abs(naive((1j, 0), (1j, 0)) - inner((1j, 0), (1j, 0))) > 1.0
print(f"PASS: 10 projector verifier checks {len(probabilities)} Born probabilities; range=[{min(probabilities):.3f},{max(probabilities):.3f}], last derived probability={probability:.3f}, max P^2-P error={max_idempotence_error:.2e}")
