/*
 * AMS "tug-of-war" F2 sketch and the bucketed (Count Sketch style) F2
 * estimator: linearity self-tests and relative error versus counter count.
 *
 *   ./ams_f2 test                         linearity / merge / turnstile checks
 *   ./ams_f2 sweep zipf|uniform TRIALS SEED
 *
 * Hash families are polynomials over GF(p), p = 2^61 - 1:
 *   degree 3 -> 4-wise independent, degree 1 -> 2-wise independent.
 * A sign is bit 0 of h(x); a bucket is h(x) mod w.  Both are off from exact
 * uniformity by O(1/p), which is far below anything measured here.
 */
#include <inttypes.h>
#include <math.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>

#define P61 ((UINT64_C(1) << 61) - 1)

static uint64_t splitmix64(uint64_t *s)
{
    uint64_t z = (*s += UINT64_C(0x9e3779b97f4a7c15));
    z = (z ^ (z >> 30)) * UINT64_C(0xbf58476d1ce4e5b9);
    z = (z ^ (z >> 27)) * UINT64_C(0x94d049bb133111eb);
    return z ^ (z >> 31);
}

static uint64_t mulmod61(uint64_t a, uint64_t b)
{
    unsigned __int128 t = (unsigned __int128)a * b;
    uint64_t lo = (uint64_t)t & P61, hi = (uint64_t)(t >> 61);
    uint64_t r = lo + hi;
    return r >= P61 ? r - P61 : r;
}

static uint64_t addmod61(uint64_t a, uint64_t b)
{
    uint64_t r = a + b;
    return r >= P61 ? r - P61 : r;
}

typedef struct { int deg; uint64_t c[4]; } poly_t;

static void poly_init(poly_t *h, int deg, uint64_t *rng)
{
    h->deg = deg;
    for (int i = 0; i <= deg; i++) h->c[i] = splitmix64(rng) % P61;
}

static uint64_t poly_eval(const poly_t *h, uint64_t x)
{
    uint64_t r = h->c[h->deg];
    for (int i = h->deg - 1; i >= 0; i--) r = addmod61(mulmod61(r, x), h->c[i]);
    return r;
}

static int64_t sign_of(const poly_t *h, uint64_t x) { return (poly_eval(h, x) & 1) ? 1 : -1; }

/* ---------- classic AMS: K independent counters Z_j = sum_i s_j(i) f_i ---------- */
typedef struct { int k; poly_t *s; int64_t *z; } ams_t;

static void ams_init(ams_t *a, int k, int deg, uint64_t seed)
{
    uint64_t rng = seed;
    a->k = k;
    a->s = malloc(sizeof *a->s * k);
    a->z = calloc(k, sizeof *a->z);
    if (!a->s || !a->z) { perror("malloc"); exit(1); }
    for (int j = 0; j < k; j++) poly_init(&a->s[j], deg, &rng);
}

static void ams_free(ams_t *a) { free(a->s); free(a->z); }

static void ams_update(ams_t *a, uint64_t x, int64_t c)
{
    for (int j = 0; j < a->k; j++) a->z[j] += sign_of(&a->s[j], x) * c;
}

/* mean of Z_j^2 over counters [lo, hi) */
static double ams_mean(const ams_t *a, int lo, int hi)
{
    double s = 0;
    for (int j = lo; j < hi; j++) s += (double)a->z[j] * (double)a->z[j];
    return s / (hi - lo);
}

/* ---------- bucketed estimator: d rows x w buckets ---------- */
typedef struct { int d, w; poly_t *b, *s; int64_t *c; } cs_t;

static void cs_init(cs_t *t, int d, int w, uint64_t seed)
{
    uint64_t rng = seed;
    t->d = d; t->w = w;
    t->b = malloc(sizeof *t->b * d);
    t->s = malloc(sizeof *t->s * d);
    t->c = calloc((size_t)d * w, sizeof *t->c);
    if (!t->b || !t->s || !t->c) { perror("malloc"); exit(1); }
    for (int r = 0; r < d; r++) { poly_init(&t->b[r], 3, &rng); poly_init(&t->s[r], 3, &rng); }
}

static void cs_free(cs_t *t) { free(t->b); free(t->s); free(t->c); }

static void cs_update(cs_t *t, uint64_t x, int64_t c)
{
    for (int r = 0; r < t->d; r++) {
        uint64_t col = poly_eval(&t->b[r], x) % (uint64_t)t->w;
        t->c[(size_t)r * t->w + col] += sign_of(&t->s[r], x) * c;
    }
}

static double cs_row(const cs_t *t, int r)
{
    double s = 0;
    for (int b = 0; b < t->w; b++) {
        double v = (double)t->c[(size_t)r * t->w + b];
        s += v * v;
    }
    return s;
}

static int cmp_double(const void *a, const void *b)
{
    double x = *(const double *)a, y = *(const double *)b;
    return (x > y) - (x < y);
}

static double median(double *v, int n)
{
    qsort(v, n, sizeof *v, cmp_double);
    return n % 2 ? v[n / 2] : 0.5 * (v[n / 2 - 1] + v[n / 2]);
}

static double cs_estimate(const cs_t *t)
{
    double r[16];
    for (int i = 0; i < t->d; i++) r[i] = cs_row(t, i);
    return median(r, t->d);
}

/* ---------- input streams ---------- */
enum { NKEYS = 10000 };

/* m draws over keys 1..NKEYS; Zipf(1.0) if zipf, else uniform */
static uint32_t *make_stream(int zipf, long m, uint64_t seed)
{
    static double cdf[NKEYS];
    double s = 0;
    for (int i = 0; i < NKEYS; i++) { s += zipf ? 1.0 / (i + 1) : 1.0; cdf[i] = s; }
    uint32_t *a = malloc(sizeof *a * m);
    if (!a) { perror("malloc"); exit(1); }
    for (long t = 0; t < m; t++) {
        double u = (splitmix64(&seed) >> 11) * 0x1.0p-53 * s;
        int lo = 0, hi = NKEYS - 1;
        while (lo < hi) { int mid = (lo + hi) / 2; if (cdf[mid] > u) hi = mid; else lo = mid + 1; }
        a[t] = (uint32_t)lo + 1;
    }
    return a;
}

static void freq_of(const uint32_t *a, long m, int64_t *f)
{
    memset(f, 0, sizeof *f * (NKEYS + 1));
    for (long t = 0; t < m; t++) f[a[t]]++;
}

static void moments(const int64_t *f, double *f2, double *f4)
{
    double s2 = 0, s4 = 0;
    for (int i = 1; i <= NKEYS; i++) { double v = (double)f[i] * f[i]; s2 += v; s4 += v * v; }
    *f2 = s2; *f4 = s4;
}

static int failures;
static void check(int ok, const char *what)
{
    printf("%s  %s\n", ok ? "ok  " : "FAIL", what);
    if (!ok) failures++;
}

static int same_ams(const ams_t *x, const ams_t *y) { return memcmp(x->z, y->z, sizeof *x->z * x->k) == 0; }
static int same_cs(const cs_t *x, const cs_t *y) { return memcmp(x->c, y->c, sizeof *x->c * x->d * x->w) == 0; }

static int run_test(void)
{
    const long m = 100000;
    const int K = 64, D = 5, W = 64;
    const uint64_t hs = 42;
    uint32_t *A = make_stream(1, m, 1), *B = make_stream(0, m, 2);
    static int64_t fa[NKEYS + 1], fb[NKEYS + 1], g[NKEYS + 1], diff[NKEYS + 1];
    freq_of(A, m, fa); freq_of(B, m, fb);

    ams_t sa, sb, sab, sf; cs_t ca, cb, cab, cf;
    ams_init(&sa, K, 3, hs); ams_init(&sb, K, 3, hs); ams_init(&sab, K, 3, hs); ams_init(&sf, K, 3, hs);
    cs_init(&ca, D, W, hs); cs_init(&cb, D, W, hs); cs_init(&cab, D, W, hs); cs_init(&cf, D, W, hs);
    for (long t = 0; t < m; t++) { ams_update(&sa, A[t], 1); cs_update(&ca, A[t], 1); }
    for (long t = 0; t < m; t++) { ams_update(&sb, B[t], 1); cs_update(&cb, B[t], 1); }
    for (long t = 0; t < m; t++) { ams_update(&sab, A[t], 1); cs_update(&cab, A[t], 1); }
    for (long t = 0; t < m; t++) { ams_update(&sab, B[t], 1); cs_update(&cab, B[t], 1); }

    for (int i = 1; i <= NKEYS; i++) if (fa[i]) { ams_update(&sf, i, fa[i]); cs_update(&cf, i, fa[i]); }
    check(same_ams(&sa, &sf) && same_cs(&ca, &cf), "sketch of stream A == sketch of its frequency vector");

    for (int j = 0; j < K; j++) sa.z[j] += sb.z[j];
    for (int j = 0; j < D * W; j++) ca.c[j] += cb.c[j];
    check(same_ams(&sa, &sab) && same_cs(&ca, &cab), "sketch(A) + sketch(B) == sketch(A then B)");

    /* strict turnstile: insert A, then delete each update of A with probability 1/2 */
    ams_t st, sg; cs_t ct, cg;
    ams_init(&st, K, 3, hs); ams_init(&sg, K, 3, hs); cs_init(&ct, D, W, hs); cs_init(&cg, D, W, hs);
    memcpy(g, fa, sizeof g);
    for (long t = 0; t < m; t++) { ams_update(&st, A[t], 1); cs_update(&ct, A[t], 1); }
    uint64_t coin = 7;
    long deleted = 0;
    for (long t = 0; t < m; t++)
        if (splitmix64(&coin) & 1) { ams_update(&st, A[t], -1); cs_update(&ct, A[t], -1); g[A[t]]--; deleted++; }
    for (int i = 1; i <= NKEYS; i++) if (g[i]) { ams_update(&sg, i, g[i]); cs_update(&cg, i, g[i]); }
    check(same_ams(&st, &sg) && same_cs(&ct, &cg), "insert A, delete half of it == sketch of the net vector");

    double f2, f4;
    moments(g, &f2, &f4);
    printf("      strict turnstile: %ld deletions, true F2 %.0f, AMS(K=%d) %.0f, CS(%dx%d) %.0f\n",
           deleted, f2, K, ams_mean(&st, 0, K), D, W, cs_estimate(&ct));

    /* general turnstile: f_A - f_B has negative entries */
    ams_t sd, sdiff; cs_t cd, cdiff;
    ams_init(&sd, K, 3, hs); ams_init(&sdiff, K, 3, hs); cs_init(&cd, D, W, hs); cs_init(&cdiff, D, W, hs);
    for (long t = 0; t < m; t++) { ams_update(&sd, A[t], 1); cs_update(&cd, A[t], 1); }
    for (long t = 0; t < m; t++) { ams_update(&sd, B[t], -1); cs_update(&cd, B[t], -1); }
    int neg = 0;
    for (int i = 1; i <= NKEYS; i++) {
        diff[i] = fa[i] - fb[i];
        neg += diff[i] < 0;
        if (diff[i]) { ams_update(&sdiff, i, diff[i]); cs_update(&cdiff, i, diff[i]); }
    }
    check(same_ams(&sd, &sdiff) && same_cs(&cd, &cdiff), "sketch(A) - sketch(B) == sketch of f_A - f_B");
    moments(diff, &f2, &f4);
    printf("      general turnstile: %d negative coordinates, true ||f_A-f_B||^2 %.0f, AMS %.0f, CS %.0f\n",
           neg, f2, ams_mean(&sd, 0, K), cs_estimate(&cd));

    ams_free(&sa); ams_free(&sb); ams_free(&sab); ams_free(&sf); ams_free(&st); ams_free(&sg);
    ams_free(&sd); ams_free(&sdiff);
    cs_free(&ca); cs_free(&cb); cs_free(&cab); cs_free(&cf); cs_free(&ct); cs_free(&cg);
    cs_free(&cd); cs_free(&cdiff);
    free(A); free(B);
    printf("%s\n", failures ? "SOME TESTS FAILED" : "all tests passed");
    return failures ? 1 : 0;
}

/* ---------- error sweep ---------- */
static const int KS[] = { 20, 40, 80, 160, 320, 640, 1280, 2560, 5120 };
enum { NK = sizeof KS / sizeof KS[0], KMAX = 5120, GROUPS = 5, NMETH = 5 };
static const char *METH[NMETH] = { "ams_mean", "ams_mom5", "cs_d1", "cs_d5", "ams_mean_2wise" };

static double quantile_abs(const double *e, int n, double q)
{
    double *v = malloc(sizeof *v * n);
    if (!v) { perror("malloc"); exit(1); }
    for (int i = 0; i < n; i++) v[i] = fabs(e[i]);
    qsort(v, n, sizeof *v, cmp_double);
    double r = v[(int)ceil(q * n) - 1];
    free(v);
    return r;
}

static int run_sweep(int zipf, int trials, uint64_t seed)
{
    const long m = 1000000;
    uint32_t *a = make_stream(zipf, m, 1);
    static int64_t f[NKEYS + 1];
    freq_of(a, m, f);
    free(a);
    double f2, f4;
    moments(f, &f2, &f4);
    double *err = malloc(sizeof *err * (size_t)NMETH * NK * trials);
    if (!err) { perror("malloc"); exit(1); }
#define ERR(meth, ki, t) err[((size_t)(meth) * NK + (ki)) * trials + (t)]

    for (int t = 0; t < trials; t++) {
        uint64_t ts = splitmix64(&seed);
        ams_t s4, s2;
        ams_init(&s4, KMAX, 3, ts);
        ams_init(&s2, KMAX, 1, ts ^ UINT64_C(0x5555));
        for (int i = 1; i <= NKEYS; i++)
            if (f[i]) { ams_update(&s4, i, f[i]); ams_update(&s2, i, f[i]); }
        for (int ki = 0; ki < NK; ki++) {
            int K = KS[ki], g = K / GROUPS;
            double means[GROUPS];
            for (int q = 0; q < GROUPS; q++) means[q] = ams_mean(&s4, q * g, (q + 1) * g);
            ERR(0, ki, t) = ams_mean(&s4, 0, K) / f2 - 1;
            ERR(1, ki, t) = median(means, GROUPS) / f2 - 1;
            ERR(4, ki, t) = ams_mean(&s2, 0, K) / f2 - 1;
            cs_t c1, c5;
            cs_init(&c1, 1, K, ts + 1 + ki);
            cs_init(&c5, GROUPS, g, ts + 101 + ki);
            for (int i = 1; i <= NKEYS; i++)
                if (f[i]) { cs_update(&c1, i, f[i]); cs_update(&c5, i, f[i]); }
            ERR(2, ki, t) = cs_estimate(&c1) / f2 - 1;
            ERR(3, ki, t) = cs_estimate(&c5) / f2 - 1;
            cs_free(&c1); cs_free(&c5);
        }
        ams_free(&s4); ams_free(&s2);
    }

    printf("# dist=%s m=%ld keys=%d trials=%d F2=%.0f F4/F2^2=%.4f\n",
           zipf ? "zipf1.0" : "uniform", m, NKEYS, trials, f2, f4 / (f2 * f2));
    printf("# theory_sd = sqrt(2*(F2^2-F4)/K)/F2; hash_per_update: ams=K, cs_d1=2, cs_d5=10\n");
    printf("method,K,theory_sd,rms,median_abs,p95_abs,mean\n");
    for (int me = 0; me < NMETH; me++)
        for (int ki = 0; ki < NK; ki++) {
            const double *e = &ERR(me, ki, 0);
            double ss = 0, s = 0;
            for (int t = 0; t < trials; t++) { ss += e[t] * e[t]; s += e[t]; }
            printf("%s,%d,%.5f,%.5f,%.5f,%.5f,%.5f\n", METH[me], KS[ki],
                   sqrt(2 * (f2 * f2 - f4) / KS[ki]) / f2, sqrt(ss / trials),
                   quantile_abs(e, trials, 0.5), quantile_abs(e, trials, 0.95), s / trials);
        }
#undef ERR
    free(err);
    return 0;
}

int main(int argc, char **argv)
{
    if (argc >= 2 && strcmp(argv[1], "test") == 0) return run_test();
    if (argc == 5 && strcmp(argv[1], "sweep") == 0)
        return run_sweep(strcmp(argv[2], "zipf") == 0, atoi(argv[3]), strtoull(argv[4], NULL, 10));
    fprintf(stderr, "usage: %s test | sweep zipf|uniform TRIALS SEED\n", argv[0]);
    return 2;
}
