# Generate perfect hash table.

def hasher(key, displacement, shift):
    KNUTH_CONSTANT = 0x9E3779B1
    return (((key ^ displacement) * KNUTH_CONSTANT) & 0xFFFFFFFF) >> shift

def build_perfect_hash(keys):
    keys_size = len(keys)
    # Round up lookup size to nearest power of two.
    lg2 = (keys_size-1).bit_length() if keys_size > 0 else 0
    lookup_size = 2**lg2
    shift = 32 - lg2

    lookup_slots = {}
    displacement_table = [0] * lookup_size
    displacement_limit = 10000

    # Hash keys into buckets (displacement=0).
    buckets = {}
    for key in keys:
        displacement_slot = hasher(key, 0, shift)
        buckets.setdefault(displacement_slot, []).append(key)

    # Order buckets largest to smallest.
    buckets_items = sorted(buckets.items(), key=lambda item: len(item[1]), reverse=True)

    for displacement_slot, bucket in buckets_items:
        bucket_size = len(bucket)
        displacement = 1
        while True:
            bucket_slots = {}
            # Hash keys at current displacement and check for collisions.
            for key in bucket:
                slot = hasher(key, displacement, shift)
                if slot in lookup_slots or slot in bucket_slots:
                    break
                bucket_slots[slot] = key

            # If no collisions, update lookup and record displacement.
            if len(bucket_slots) == bucket_size:
                lookup_slots.update(bucket_slots)
                displacement_table[displacement_slot] = displacement
                break

            displacement += 1
            if displacement >= displacement_limit:
                raise RuntimeError("Unable to build perfect hash")

    # Map slots dictionary to table array.
    lookup_table = [None] * lookup_size
    for slot, key in lookup_slots.items():
        lookup_table[slot] = key

    return lookup_table, displacement_table, shift

# Main.

from random import sample

keys = list(sample(range(10000), k=8))
lookup_table, displacement_table, shift = build_perfect_hash(keys)
lookup_size = len(lookup_table)

print(keys)
# print(lookup_size)
# print(displacement_table)
# print(shift)

def perfect_hash(key):
    displacement_slot = hasher(key, 0, shift)
    slot = hasher(key, displacement_table[displacement_slot], shift)
    return slot

def perfect_lookup(key):
    return lookup_table[perfect_hash(key)]

for i, displacement in enumerate(displacement_table):
    if displacement != 0:
        print(f"Slot: {i} -> Displacement {displacement}")
for key in keys:
    slot = perfect_hash(key)
    print(f"Key: {key} -> To unique slot {slot} (Stored {lookup_table[slot]})")