/* test_scan.c -- differential tests (SIMD vs scalar) and page-end tests. */
#define _GNU_SOURCE
#include <signal.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <sys/mman.h>
#include <sys/wait.h>
#include <unistd.h>
#include "scan.h"

static uint64_t rng = 0x9E3779B97F4A7C15ULL;
static uint64_t next_rand(void)
{
    rng ^= rng << 13; rng ^= rng >> 7; rng ^= rng << 17;
    return rng;
}

static long failures;
#define CHECK(cond, ...) do { if (!(cond)) { failures++; \
    if (failures < 20) { printf("FAIL %s:%d: ", __FILE__, __LINE__); \
    printf(__VA_ARGS__); printf("\n"); } } } while (0)

static void random_fill(char *b, size_t n, const char *alpha)
{
    size_t k = strlen(alpha);
    for (size_t i = 0; i < n; i++)
        b[i] = alpha[next_rand() % k];
}

static void test_memchr(void)
{
    static char buf[512 + 64];
    long cases = 0;
    for (size_t off = 0; off < 64; off++)
        for (size_t n = 0; n <= 512; n++) {
            char *s = buf + off;
            for (size_t i = 0; i < n; i++)
                s[i] = (char)(next_rand() % 256);
            int c = (int)(next_rand() % 256);
            if (n && next_rand() % 2) s[next_rand() % n] = (char)c;
            const char *r0 = memchr_scalar(s, c, n);
            CHECK(memchr_sse2(s, c, n) == r0, "memchr_sse2 off=%zu n=%zu", off, n);
            CHECK(memchr_avx2(s, c, n) == r0, "memchr_avx2 off=%zu n=%zu", off, n);
            CHECK(memchr(s, c, n) == r0, "glibc memchr off=%zu n=%zu", off, n);
            cases++;
        }
    printf("memchr: %ld cases (64 offsets x n=0..512)\n", cases);
}

static void test_strlen_find(void)
{
    static char buf[400];
    long cases = 0;
    const char set[4] = { ',', '\n', '"', '\r' };
    for (size_t off = 0; off < 64; off++)
        for (size_t len = 0; len <= 300; len++) {
            char *s = buf + off;
            for (size_t i = 0; i < len; i++)
                s[i] = (char)(1 + next_rand() % 255);
            s[len] = 0;
            CHECK(strlen_avx2_aligned(s) == len, "strlen_aligned off=%zu len=%zu", off, len);
            for (size_t i = 0; i < len; i++)
                if (s[i] == set[0] || s[i] == set[1] || s[i] == set[2] || s[i] == set[3])
                    s[i] = 'x';
            if (len && next_rand() % 2) s[next_rand() % len] = set[next_rand() % 4];
            size_t r0 = find_any4_scalar(s, len, set);
            CHECK(find_any4_pcmpistri(s, len, set) == r0, "pcmpistri len=%zu", len);
            CHECK(find_any4_avx2(s, len, set) == r0, "any4_avx2 len=%zu", len);
            cases++;
        }
    printf("strlen/find_any4: %ld cases (64 offsets x len=0..300)\n", cases);
}

static uint32_t out0[4096], out1[4096];

static int same_marks(size_t c0, size_t c1, int o0, int o1)
{
    return c0 == c1 && o0 == o1 && memcmp(out0, out1, c0 * sizeof out0[0]) == 0;
}

static void test_csv_json(void)
{
    static char buf[4096];
    const char *csv_alpha = "\"\",,\naab";
    const char *json_alpha = "\"\\\\\\{}[]:,a ";
    long cases = 0;
    for (int iter = 0; iter < 20000; iter++) {
        size_t n = next_rand() % 1500;
        int o0, o1;
        random_fill(buf, n, csv_alpha);
        size_t c0 = csv_seps_scalar(buf, n, out0, &o0);
        size_t c1 = csv_seps_simd(buf, n, out1, &o1);
        CHECK(same_marks(c0, c1, o0, o1), "csv iter=%d n=%zu", iter, n);
        random_fill(buf, n, json_alpha);
        c0 = json_marks_scalar(buf, n, out0, &o0);
        c1 = json_marks_simd(buf, n, out1, &o1);
        CHECK(same_marks(c0, c1, o0, o1), "json iter=%d n=%zu", iter, n);
        CHECK(json_count_simd(buf, n, out1, &o1) == c0 && o1 == o0, "json count iter=%d", iter);
        cases += 2;
    }
    for (int iter = 0; iter < 100000; iter++) {
        uint64_t x = next_rand() & next_rand(), ref = 0, acc = 0;
        for (int i = 0; i < 64; i++) {
            acc ^= (x >> i) & 1;
            ref |= acc << i;
        }
        CHECK(prefix_xor_clmul(x) == ref, "prefix_xor %016llx", (unsigned long long)x);
    }
    printf("csv/json: %ld random inputs (n=0..1499); prefix_xor: 100000 words\n", cases);
}

/* One readable page followed by a PROT_NONE guard page. */
static char *page_with_guard(void)
{
    long pg = sysconf(_SC_PAGESIZE);
    char *m = mmap(NULL, 2 * (size_t)pg, PROT_READ | PROT_WRITE,
                   MAP_PRIVATE | MAP_ANONYMOUS, -1, 0);
    if (m == MAP_FAILED || mprotect(m + pg, (size_t)pg, PROT_NONE) != 0) {
        perror("mmap/mprotect");
        exit(2);
    }
    memset(m, 'a', (size_t)pg);
    return m + pg; /* first byte of the guard page */
}

/* Run f(s) in a child; return the signal that killed it, or 0. */
static int crashes(size_t (*f)(const char *), const char *s)
{
    fflush(stdout);
    pid_t pid = fork();
    if (pid == 0) {
        volatile size_t r = f(s);
        (void)r;
        _exit(0);
    }
    int st;
    waitpid(pid, &st, 0);
    return WIFSIGNALED(st) ? WTERMSIG(st) : 0;
}

static void test_page_end(void)
{
    char *end = page_with_guard();
    int naive_segv = 0, first_ok = -1;
    long cases = 0;
    for (size_t len = 0; len <= 200; len++) {
        char *s = end - len - 1; /* string + NUL ends exactly at the guard */
        memset(s, 'a', len);
        s[len] = 0;
        CHECK(strlen_avx2_aligned(s) == len, "page-end strlen_aligned len=%zu", len);
        const char *r0 = memchr_scalar(s, 'b', len);
        CHECK(memchr_sse2(s, 'b', len) == r0, "page-end memchr_sse2 len=%zu", len);
        CHECK(memchr_avx2(s, 'b', len) == r0, "page-end memchr_avx2 len=%zu", len);
        int o0, o1;
        size_t c0 = csv_seps_scalar(s, len, out0, &o0);
        size_t c1 = csv_seps_simd(s, len, out1, &o1);
        CHECK(same_marks(c0, c1, o0, o1), "page-end csv len=%zu", len);
        c0 = json_marks_scalar(s, len, out0, &o0);
        c1 = json_marks_simd(s, len, out1, &o1);
        CHECK(same_marks(c0, c1, o0, o1), "page-end json len=%zu", len);
        int sig = crashes(strlen_avx2_naive, s);
        if (sig == SIGSEGV) naive_segv++;
        else if (sig == 0 && first_ok < 0) first_ok = (int)len;
        cases++;
    }
    printf("page-end: %ld lengths (0..200), string ends at a PROT_NONE page\n", cases);
    printf("  strlen_avx2_naive: SIGSEGV for %d lengths; first length without fault: %d\n",
           naive_segv, first_ok);
}

int main(void)
{
    test_memchr();
    test_strlen_find();
    test_csv_json();
    test_page_end();
    printf("%s (%ld failures)\n", failures ? "FAILED" : "ALL PASSED", failures);
    return failures != 0;
}
