const std = @import("std");
const builtin = @import("builtin");
const print = std.debug.print;

const use_bmi2 = builtin.cpu.arch == .x86_64;

const WORDS = 100000;
const PASSES = 1000;

fn L(comptime T: type, comptime k: u6) T {
    comptime {
        var acc: T = 0;
        var i = 0;
        while (i < @bitSizeOf(T)) : (i += k) acc |= 1 << i;
        return acc;
    }
}

fn H(comptime T: type, comptime k: u6) T {
    return L(T, k) << (k - 1);
}

fn leX(comptime k: u6, x: u64, y: u64) u64 {
    const hk = comptime H(u64, k);
    return (((y | hk) - (x & ~hk)) ^ x ^ y) & hk;
}

fn gtX0(comptime k: u6, x: u64) u64 {
    const hk = comptime H(u64, k);
    return (((x | hk) - comptime L(u64, k)) | x) & hk;
}


fn selectNaive(x: u64, r: u6) !u6 {
    var bits = x;
    var remaining: u64 = r;

    for (0..64) |i| {
        if ((bits & 1) != 0) {
            if (remaining == 0) return @intCast(i);
            remaining -= 1;
        }
        bits >>= 1;
    }

    return error.NotFound;
}

fn selectSkip(x: u64, r: u6) !u6 {
    var v = x;
    for (0..r) |_| v &= v - 1;

    const pos = @ctz(v);
    if (pos == 64) return error.NotFound;
    return @intCast(pos);
}

fn selectBroadword(x: u64, r: u6) !u6 {
    const L8: u64 = comptime L(u64, 8);

    var s = x - ((x & 0xAAAAAAAAAAAAAAAA) >> 1);
    s = (s & 0x3333333333333333) + ((s >> 2) & 0x3333333333333333);
    s = ((s + (s >> 4)) & 0x0F0F0F0F0F0F0F0F) *% L8;

    var b = ((leX(8, s, r *% L8) >> 7) *% L8 >> 53);
    const l = r - ((std.math.shr(u64, s << 8, b)) & 0xFF);

    s = (gtX0(8, ((std.math.shr(u64, x, b) & 0xFF) *% L8) & 0x8040201008040201) >> 7) *% L8;
    b += ((leX(8, s, l *% L8) >> 7) *% L8 >> 56);

    if (b == 72) return error.NotFound;
    return @intCast(b);
}

fn selectBMI2(x: u64, r: u6) !u6 {
    const mask = @as(u64, 1) << r;
    const deposited = asm ("pdep %[x], %[mask], %[result]"
        : [result] "=r" (-> u64),
        : [mask] "r" (mask),
          [x] "r" (x),
    );
    if (deposited == 0) return error.NotFound;
    return @intCast(@ctz(deposited));
}

const Density = enum {
    sparse,
    random,
    dense,

    fn nextWord(density: Density, rng: std.Random) u64 {
        return switch (density) {
            .sparse => rng.int(u64) & rng.int(u64) & rng.int(u64) & rng.int(u64),
            .random => rng.int(u64),
            .dense => rng.int(u64) | rng.int(u64) | rng.int(u64) | rng.int(u64),
        };
    }
};

fn fill(rng: std.Random, density: Density, words: []u64, ranks: []u6) f64 {
    var total_bits: u64 = 0;

    for (words, ranks) |*word, *rank| {
        word.* = density.nextWord(rng);
        while (word.* == 0) word.* = density.nextWord(rng);

        rank.* = @intCast(rng.uintLessThan(u64, @popCount(word.*)));
        total_bits += @popCount(word.*);
    }

    return @as(f64, @floatFromInt(total_bits)) / @as(f64, @floatFromInt(words.len));
}

fn verify(words: []const u64, ranks: []const u6) !void {
    for (words, ranks) |word, rank| {
        const expected = try selectNaive(word, rank);
        if (try selectSkip(word, rank) != expected) return error.SkipMismatch;
        if (try selectBroadword(word, rank) != expected) return error.BroadwordMismatch;
        if (comptime use_bmi2) {
            if (try selectBMI2(word, rank) != expected) return error.BMI2Mismatch;
        }
    }
}

fn bench(io: std.Io, comptime select: anytype, words: []const u64, ranks: []const u6) f64 {
    var sum: u64 = 0;
    for (words, ranks) |word, rank| sum +%= select(word, rank) catch unreachable;

    const start = std.Io.Clock.awake.now(io);
    for (0..PASSES) |_| {
        for (words, ranks) |word, rank| sum +%= select(word, rank) catch unreachable;
    }
    const elapsed = start.durationTo(std.Io.Clock.awake.now(io));

    std.mem.doNotOptimizeAway(sum);
    return @as(f64, @floatFromInt(elapsed.toNanoseconds())) / @as(f64, @floatFromInt(words.len * PASSES));
}

pub fn main(init: std.process.Init) !void {
    const io = init.io;
    const allocator = std.heap.page_allocator;

    const words = try allocator.alloc(u64, WORDS);
    defer allocator.free(words);
    const ranks = try allocator.alloc(u6, WORDS);
    defer allocator.free(ranks);

    var prng = std.Random.DefaultPrng.init(0xFEEDBABE);
    const rng = prng.random();

    print("{d} words, {d} passes, ns per call\n", .{ WORDS, PASSES });
    print("{s:<24} {s:>8} {s:>8} {s:>11}", .{ "input", "naive", "skip", "broadword" });
    if (comptime use_bmi2) print(" {s:>8}", .{"bmi2"});
    print("\n", .{});

    for (std.enums.values(Density)) |density| {
        const avg_bits = fill(rng, density, words, ranks);
        try verify(words, ranks);

        var buf: [24]u8 = undefined;
        const label = try std.fmt.bufPrint(&buf, "{s}, ~{d:.0} set bits", .{ @tagName(density), avg_bits });

        print("{s:<24} {d:>8.1} {d:>8.1} {d:>11.1}", .{
            label,
            bench(io, selectNaive, words, ranks),
            bench(io, selectSkip, words, ranks),
            bench(io, selectBroadword, words, ranks),
        });
        if (comptime use_bmi2) print(" {d:>8.1}", .{bench(io, selectBMI2, words, ranks)});
        print("\n", .{});
    }
}
