/* Brute-force checks of every guarantee used in the article. */
#include "summaries.h"

#include <stdio.h>
#include <stdlib.h>
#include <string.h>

static uint64_t rng = 0x9e3779b97f4a7c15ULL;
static uint64_t next(void)
{
    uint64_t z = (rng += 0x9e3779b97f4a7c15ULL);
    z = (z ^ (z >> 30)) * 0xbf58476d1ce4e5b9ULL;
    z = (z ^ (z >> 27)) * 0x94d049bb133111ebULL;
    return z ^ (z >> 31);
}

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

/* stream kinds: 0 uniform, 1 skewed, 2 round-robin over k+1 items, 3 sorted runs */
static int make_stream(int kind, int u, int n, int k, uint64_t *out)
{
    for (int i = 0; i < n; i++) {
        uint64_t r = next();
        switch (kind) {
        case 0: out[i] = r % (uint64_t)u; break;
        case 1: { double p = (double)(r >> 11) / 9007199254740992.0; out[i] = (uint64_t)((double)u * p * p * p); break; }
        case 2: out[i] = (uint64_t)(i % (k + 1)); break;
        default: out[i] = (uint64_t)((i * (int64_t)u) / n); break;
        }
    }
    return n;
}

static void check_one(const uint64_t *st, int n, int u, int k, int64_t *f)
{
    memset(f, 0, sizeof(int64_t) * (size_t)u);
    mg_t mg; ss_t ss, ssk; lc_t lc;
    mg_init(&mg, k); ss_init(&ss, k + 1); ss_init(&ssk, k); lc_init(&lc, k);
    for (int i = 0; i < n; i++) {
        f[st[i]]++;
        mg_update(&mg, st[i]); ss_update(&ss, st[i]); ss_update(&ssk, st[i]); lc_update(&lc, st[i]);
    }
    int64_t nhat = mg_sum(&mg), B = mg_bound(&mg);
    CHECK((n - nhat) % (k + 1) == 0, "n-nhat=%lld not divisible by k+1=%d", (long long)(n - nhat), k + 1);
    CHECK(mg.size <= k, "mg size %d > k %d", mg.size, k);
    CHECK((n - nhat) / (k + 1) == mg.decrements, "decrement count");
    CHECK(ss_min(&ss) == B, "isomorphism: min_SS=%lld, (n-nhat)/(k+1)=%lld", (long long)ss_min(&ss), (long long)B);
    CHECK(ss_min(&ssk) <= n / k, "min <= n/m");
    for (int x = 0; x < u; x++) {
        int64_t e = mg_estimate(&mg, (uint64_t)x);
        CHECK(e <= f[x] && f[x] <= e + B, "MG x=%d f=%lld est=%lld B=%lld", x, (long long)f[x], (long long)e, (long long)B);
        CHECK(ss_count(&ss, (uint64_t)x) - e == ss_min(&ss), "isomorphism x=%d", x);
        int64_t c = ss_count(&ssk, (uint64_t)x), er = ss_err(&ssk, (uint64_t)x);
        if (er >= 0) {
            CHECK(c - er <= f[x] && f[x] <= c && er <= ss_min(&ssk), "SS tracked x=%d", x);
        } else {
            CHECK(f[x] <= ss_min(&ssk), "SS untracked x=%d f=%lld min=%lld", x, (long long)f[x], (long long)ss_min(&ssk));
        }
        CHECK(lc_lower(&lc, (uint64_t)x) <= f[x] && f[x] <= lc_upper(&lc, (uint64_t)x), "LC x=%d", x);
        CHECK(lc_upper(&lc, (uint64_t)x) - lc_lower(&lc, (uint64_t)x) <= n / k, "LC width x=%d", x);
    }
    /* SS keeps exact stream length in its counters once full */
    if (ssk.used == k) {
        uint64_t *kk = malloc(sizeof(uint64_t) * (size_t)k);
        int64_t *cc = malloc(sizeof(int64_t) * (size_t)k), *ee = malloc(sizeof(int64_t) * (size_t)k);
        int got = ss_topk(&ssk, kk, cc, ee, k);
        int64_t sum = 0;
        for (int i = 0; i < got; i++) { sum += cc[i]; if (i) CHECK(cc[i - 1] >= cc[i], "topk order"); }
        CHECK(got == k && sum == n, "SS sum of counters %lld != n %d", (long long)sum, n);
        free(kk); free(cc); free(ee);
    }
    mg_free(&mg); ss_free(&ss); ss_free(&ssk); lc_free(&lc);
}

/* split into P parts, summarise each, merge pairwise in a tree */
static void check_merge(const uint64_t *st, int n, int u, int k, const int64_t *f, int P)
{
    mg_t part[16];
    for (int p = 0; p < P; p++) {
        mg_init(&part[p], k);
        for (int i = p * n / P; i < (p + 1) * n / P; i++) mg_update(&part[p], st[i]);
    }
    for (int w = 1; w < P; w *= 2) {
        for (int p = 0; p + w < P; p += 2 * w) {
            mg_t out;
            mg_merge(&out, &part[p], &part[p + w]);
            mg_free(&part[p]); mg_free(&part[p + w]);
            part[p] = out;
            mg_init(&part[p + w], k);
        }
    }
    int64_t B = mg_bound(&part[0]);
    CHECK(part[0].n == n && part[0].size <= k, "merged size");
    for (int x = 0; x < u; x++) {
        int64_t e = mg_estimate(&part[0], (uint64_t)x);
        CHECK(e <= f[x] && f[x] <= e + B, "merge x=%d f=%lld est=%lld B=%lld", x, (long long)f[x], (long long)e, (long long)B);
    }
    for (int p = 0; p < P; p++) mg_free(&part[p]);
}

int main(void)
{
    enum { NMAX = 20000, UMAX = 3000 };
    uint64_t *st = malloc(sizeof(uint64_t) * NMAX);
    int64_t *f = malloc(sizeof(int64_t) * UMAX);
    int trials = 0;
    for (int rep = 0; rep < 60; rep++) {
        for (int kind = 0; kind < 4; kind++) {
            int k = 1 + (int)(next() % 60);
            int u = 2 + (int)(next() % (UMAX - 2));
            if (kind == 2 && u < k + 1) u = k + 1;
            int n = 1 + (int)(next() % NMAX);
            make_stream(kind, u, n, k, st);
            check_one(st, n, u, k, f);
            check_merge(st, n, u, k, f, 2 + (int)(next() % 15));
            trials++;
        }
    }
    printf("%d random streams, %ld checks, %ld failures\n", trials, checks, failures);
    free(st); free(f);
    return failures != 0;
}
