#!/usr/bin/env python3
"""
SHA-256 truncated-collision search, A100, packed single-slot table.

Difference vs. collide_a100.py: the two uint64 tables (htab + ctab) are replaced
by ONE uint64 table, so each hash costs a single random atomicExch instead of
three random accesses.  The measured ceiling with an L2-resident table was
5.45 GH/s, so this should land close to it.

Slot layout (64 bits):
    bit  63     : occupied flag
    bits 62..39 : 24-bit tag  = hash bits [table_log .. table_log+23]
    bits 38..0  : 39-bit counter offset relative to the epoch base

Because only 24 tag bits are compared on-GPU (plus the table_log index bits
that are equal by construction), roughly 1 in 2^24 hashes reports a false
candidate.  Those are cheap: the host re-hashes both messages and drops them,
and they show up in the 'fp' counter.  True collisions still require a full
match on all `--bits` bits.

The 39-bit offset field means the table must be wiped and the epoch rebased
every 2^39 hashes (~100 s at 5 GH/s).  The wipe costs ~15 ms and the table
refills in well under a second, so the duty-cycle loss is under 1%.

    python collide_a100_packed.py --prefix HackQuest2025 --bits 64
"""

import argparse
import hashlib
import json
import math
import random
import struct
import time

import cupy as cp
import numpy as np

ASCII_BASE = 95
MSG_LEN = 16

OFFSET_BITS = 39
TAG_BITS = 24
OFFSET_MASK = (1 << OFFSET_BITS) - 1
TAG_MASK = (1 << TAG_BITS) - 1
EPOCH_SPAN = 1 << OFFSET_BITS

KERNEL_SRC = r'''
extern "C" {

#define ASCII_BASE  95
#define MSG_LEN     __MSG_LEN__
#define PREFIX_LEN  __PREFIX_LEN__
#define TPB         __TPB__

#define OFFSET_BITS 39
#define OFFSET_MASK ((1ULL << OFFSET_BITS) - 1ULL)
#define TAG_MASK    0xFFFFFFULL
#define OCCUPIED    (1ULL << 63)

__constant__ unsigned int k[64] = {
    0x428a2f98, 0x71374491, 0xb5c0fbcf, 0xe9b5dba5, 0x3956c25b, 0x59f111f1, 0x923f82a4, 0xab1c5ed5,
    0xd807aa98, 0x12835b01, 0x243185be, 0x550c7dc3, 0x72be5d74, 0x80deb1fe, 0x9bdc06a7, 0xc19bf174,
    0xe49b69c1, 0xefbe4786, 0x0fc19dc6, 0x240ca1cc, 0x2de92c6f, 0x4a7484aa, 0x5cb0a9dc, 0x76f988da,
    0x983e5152, 0xa831c66d, 0xb00327c8, 0xbf597fc7, 0xc6e00bf3, 0xd5a79147, 0x06ca6351, 0x14292967,
    0x27b70a85, 0x2e1b2138, 0x4d2c6dfc, 0x53380d13, 0x650a7354, 0x766a0abb, 0x81c2c92e, 0x92722c85,
    0xa2bfe8a1, 0xa81a664b, 0xc24b8b70, 0xc76c51a3, 0xd192e819, 0xd6990624, 0xf40e3585, 0x106aa070,
    0x19a4c116, 0x1e376c08, 0x2748774c, 0x34b0bcb5, 0x391c0cb3, 0x4ed8aa4a, 0x5b9cca4f, 0x682e6ff3,
    0x748f82ee, 0x78a5636f, 0x84c87814, 0x8cc70208, 0x90befffa, 0xa4506ceb, 0xbef9a3f7, 0xc67178f2
};

__device__ __forceinline__ unsigned int rotr(unsigned int x, unsigned int n) {
    return (x >> n) | (x << (32 - n));
}
__device__ __forceinline__ unsigned int ch(unsigned int x, unsigned int y, unsigned int z) {
    return z ^ (x & (y ^ z));
}
__device__ __forceinline__ unsigned int maj(unsigned int x, unsigned int y, unsigned int z) {
    return (x & y) | (z & (x | y));
}
__device__ __forceinline__ unsigned int S0(unsigned int x) { return rotr(x,2) ^ rotr(x,13) ^ rotr(x,22); }
__device__ __forceinline__ unsigned int S1(unsigned int x) { return rotr(x,6) ^ rotr(x,11) ^ rotr(x,25); }
__device__ __forceinline__ unsigned int s0(unsigned int x) { return rotr(x,7) ^ rotr(x,18) ^ (x >> 3); }
__device__ __forceinline__ unsigned int s1(unsigned int x) { return rotr(x,17) ^ rotr(x,19) ^ (x >> 10); }

__device__ __forceinline__ unsigned long long sha256_top64(unsigned int *W) {
    unsigned int a = 0x6a09e667, b = 0xbb67ae85, c = 0x3c6ef372, d = 0xa54ff53a;
    unsigned int e = 0x510e527f, f = 0x9b05688c, g = 0x1f83d9ab, h = 0x5be0cd19;

    #pragma unroll
    for (int i = 0; i < 64; ++i) {
        unsigned int wi;
        if (i < 16) {
            wi = W[i];
        } else {
            wi = W[i & 15] = W[i & 15] + s0(W[(i - 15) & 15])
                           + W[(i - 7) & 15] + s1(W[(i - 2) & 15]);
        }
        unsigned int t1 = h + S1(e) + ch(e, f, g) + k[i] + wi;
        unsigned int t2 = S0(a) + maj(a, b, c);
        h = g; g = f; f = e; e = d + t1;
        d = c; c = b; b = a; a = t1 + t2;
    }
    unsigned int h0 = 0x6a09e667u + a;
    unsigned int h1 = 0xbb67ae85u + b;
    return ((unsigned long long)h0 << 32) | (unsigned long long)h1;
}

__global__ void __launch_bounds__(TPB, 4)
search_kernel(const unsigned int * __restrict__ wbase,
              unsigned long long hi_base,
              unsigned long long epoch_base,
              int shift,                       /* 64 - collision_bits */
              int table_log,
              unsigned long long table_mask,
              unsigned long long * __restrict__ tab,
              unsigned long long * __restrict__ results,
              unsigned int * __restrict__ result_count,
              unsigned int result_cap,
              int inner_rounds)
{
    const unsigned long long gid   = (unsigned long long)blockIdx.x * blockDim.x + threadIdx.x;
    const unsigned long long total = (unsigned long long)gridDim.x * blockDim.x;

    unsigned int wb[16];
    #pragma unroll
    for (int i = 0; i < 16; ++i) wb[i] = wbase[i];

    for (int r = 0; r < inner_rounds; ++r) {
        const unsigned long long hi = hi_base + gid + (unsigned long long)r * total;

        unsigned int wfix[16];
        #pragma unroll
        for (int i = 0; i < 16; ++i) wfix[i] = wb[i];

        unsigned long long c = hi;
        #pragma unroll
        for (int i = 1; i < MSG_LEN; ++i) {
            unsigned int dgt = (unsigned int)(c % ASCII_BASE);
            c /= ASCII_BASE;
            const int p = PREFIX_LEN + i;
            wfix[p >> 2] |= (0x20u + dgt) << (24 - 8 * (p & 3));
        }

        const int wi0 = PREFIX_LEN >> 2;
        const int sh0 = 24 - 8 * (PREFIX_LEN & 3);
        const unsigned long long cnt_base = hi * ASCII_BASE;

        for (unsigned int j = 0; j < ASCII_BASE; ++j) {
            unsigned int W[16];
            #pragma unroll
            for (int i = 0; i < 16; ++i) W[i] = wfix[i];
            W[wi0] |= (0x20u + j) << sh0;

            const unsigned long long hp  = sha256_top64(W) >> shift;
            const unsigned long long idx = hp & table_mask;
            const unsigned long long tag = (hp >> table_log) & TAG_MASK;
            const unsigned long long off = (cnt_base + j) - epoch_base;

            const unsigned long long packed = OCCUPIED | (tag << OFFSET_BITS) | off;
            const unsigned long long old    = atomicExch(&tab[idx], packed);

            if (old != 0ULL && ((old >> OFFSET_BITS) & TAG_MASK) == tag) {
                const unsigned long long ooff = old & OFFSET_MASK;
                if (ooff != off) {
                    unsigned int slot = atomicAdd(result_count, 1u);
                    if (slot < result_cap) {
                        results[2 * slot + 0] = ooff;
                        results[2 * slot + 1] = off;
                    }
                }
            }
        }
    }
}

}  /* extern "C" */
'''


def counter_to_ascii(cnt, msg_len=MSG_LEN):
    out = []
    for _ in range(msg_len):
        out.append(chr(0x20 + int(cnt % ASCII_BASE)))
        cnt //= ASCII_BASE
    return "".join(out)


def build_wbase(prefix_bytes, msg_len=MSG_LEN):
    mlen = len(prefix_bytes) + msg_len
    if mlen + 9 > 64:
        raise ValueError("prefix + suffix must fit in one 64-byte block")
    blk = bytearray(64)
    blk[:len(prefix_bytes)] = prefix_bytes
    blk[mlen] = 0x80
    struct.pack_into(">Q", blk, 56, mlen * 8)
    return np.frombuffer(bytes(blk), dtype=">u4").astype(np.uint32)


def pick_table_log(requested=None, reserve_frac=0.75):
    """One uint64 table => 8 bytes per slot."""
    if requested:
        return requested
    free, _ = cp.cuda.runtime.memGetInfo()
    log = int(math.floor(math.log2(free * reserve_frac / 8.0)))
    return max(20, min(31, log))


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--prefix", default="HackQuest2025")
    ap.add_argument("--bits", type=int, default=64)
    ap.add_argument("--threads", type=int, default=256)
    ap.add_argument("--blocks-per-sm", type=int, default=8)
    ap.add_argument("--inner", type=int, default=8)
    ap.add_argument("--table-log", type=int, default=0)
    ap.add_argument("--result-cap", type=int, default=16384)
    ap.add_argument("--out", default="collisions.jsonl")
    args = ap.parse_args()

    if not (1 <= args.bits <= 64):
        raise SystemExit("--bits must be in 1..64")

    prefix_bytes = args.prefix.encode("ascii")
    props = cp.cuda.runtime.getDeviceProperties(0)
    sm_count = props["multiProcessorCount"]
    name = props["name"]
    name = name.decode() if isinstance(name, bytes) else name

    tpb = args.threads
    blocks = sm_count * args.blocks_per_sm
    total_threads = tpb * blocks
    hashes_per_launch = total_threads * args.inner * ASCII_BASE
    if hashes_per_launch >= EPOCH_SPAN:
        raise SystemExit("one launch exceeds the 2^39 epoch span; lower --inner")

    table_log = pick_table_log(args.table_log or None)
    table_slots = 1 << table_log

    src = (KERNEL_SRC
           .replace("__MSG_LEN__", str(MSG_LEN))
           .replace("__PREFIX_LEN__", str(len(prefix_bytes)))
           .replace("__TPB__", str(tpb)))
    kernel = cp.RawKernel(src, "search_kernel", options=("--std=c++11",))

    wbase_gpu = cp.asarray(build_wbase(prefix_bytes))
    tab = cp.zeros(table_slots, dtype=cp.uint64)
    results = cp.zeros(2 * args.result_cap, dtype=cp.uint64)
    result_count = cp.zeros(1, dtype=cp.uint32)

    true_rate_log = args.bits - table_log
    print("--- SHA-256 truncated collision search (packed table) ---")
    print(f"GPU              : {name} ({sm_count} SMs)")
    print(f"Prefix           : {args.prefix!r} ({len(prefix_bytes)} bytes)")
    print(f"Collision bits   : {args.bits}")
    print(f"Grid             : {blocks} blocks x {tpb} threads = {total_threads:,} threads")
    print(f"Hashes / launch  : {hashes_per_launch:,}")
    print(f"Table            : 2^{table_log} slots ({table_slots * 8 / 2**30:.1f} GiB, 1 atomic/hash)")
    print(f"Expected hashes / collision : ~2^{true_rate_log}")
    print(f"Expected false candidates   : ~1 per 2^{TAG_BITS} hashes")
    print(f"Epoch span       : 2^{OFFSET_BITS} hashes between table wipes")
    print(f"Output           : {args.out}")
    print()

    shift = 64 - args.bits
    hi_base = random.randint(1, 10**6)
    epoch_base = hi_base * ASCII_BASE
    total_hashes = 0
    found = dupes = false_pos = overflow = epochs = 0
    seen = set()
    t0 = time.time()

    outf = open(args.out, "a", buffering=1)
    try:
        while True:
            counter_start = hi_base * ASCII_BASE
            if counter_start + hashes_per_launch - epoch_base >= EPOCH_SPAN:
                tab.fill(0)
                epoch_base = counter_start
                epochs += 1

            kernel((blocks,), (tpb,), (
                wbase_gpu, np.uint64(hi_base), np.uint64(epoch_base),
                np.int32(shift), np.int32(table_log), np.uint64(table_slots - 1),
                tab, results, result_count, np.uint32(args.result_cap),
                np.int32(args.inner),
            ))

            n = int(result_count.get()[0])
            if n:
                if n > args.result_cap:
                    overflow += n - args.result_cap
                take = min(n, args.result_cap)
                data = results[:2 * take].get().reshape(take, 2)
                result_count.fill(0)

                for o1, o2 in data:
                    c1 = epoch_base + int(o1)
                    c2 = epoch_base + int(o2)
                    key = (min(c1, c2), max(c1, c2))
                    if key in seen:
                        dupes += 1
                        continue

                    m1 = args.prefix + counter_to_ascii(c1)
                    m2 = args.prefix + counter_to_ascii(c2)
                    h1 = hashlib.sha256(m1.encode("ascii")).digest()
                    h2 = hashlib.sha256(m2.encode("ascii")).digest()
                    d1 = int.from_bytes(h1[:8], "big")
                    d2 = int.from_bytes(h2[:8], "big")
                    if (d1 >> shift) != (d2 >> shift):
                        false_pos += 1
                        continue

                    seen.add(key)
                    found += 1
                    rec = {
                        "n": found,
                        "bits": args.bits,
                        "msg1": m1,
                        "msg2": m2,
                        "counter1": c1,
                        "counter2": c2,
                        "prefix_hex": hex(d1 >> shift),
                        "sha1": h1.hex() if False else hashlib.sha256(m1.encode("ascii")).hexdigest(),
                        "sha2": hashlib.sha256(m2.encode("ascii")).hexdigest(),
                        "hashes": total_hashes,
                        "elapsed": round(time.time() - t0, 2),
                    }
                    outf.write(json.dumps(rec, ensure_ascii=False) + "\n")
                    print(f"\n[{found}] collision @ {hex(d1 >> shift)}"
                          f"  ({total_hashes:,} hashes, {rec['elapsed']}s)")
                    print(f"    |{m1}|")
                    print(f"    |{m2}|")

            total_hashes += hashes_per_launch
            hi_base += total_threads * args.inner

            el = time.time() - t0
            rate = total_hashes / el if el > 0 else 0
            print(f"{total_hashes:18,d} hashes | {rate/1e9:6.3f} GH/s | "
                  f"found {found} | dup {dupes} | fp {false_pos} | "
                  f"ep {epochs}{' | OVF ' + str(overflow) if overflow else ''}",
                  end="\r", flush=True)

    except KeyboardInterrupt:
        el = time.time() - t0
        print(f"\n\nStopped. {total_hashes:,} hashes in {el:.1f}s "
              f"({total_hashes/el/1e9:.3f} GH/s), {found} collisions -> {args.out}")
        if overflow:
            print(f"warning: {overflow} candidates dropped; raise --result-cap")
    finally:
        outf.close()


if __name__ == "__main__":
    main()
