/*
 * runs.c -- run generation for external sorting:
 *   load-sort-store (fill M records, quicksort, write)  vs  replacement selection.
 *
 * Reports run counts / run lengths (in units of M), key comparisons per record,
 * and (optionally) CPU time.  Everything is in memory: "writing a run" just
 * records its length, so the program measures the algorithms, not a disk.
 *
 *   gcc -O2 -Wall -Wextra -o runs runs.c
 *   ./runs trace          # step-by-step trace of replacement selection, M = 3
 *   ./runs lengths        # run-length table, M = 65536, N = 256 M, 3 seeds
 *   ./runs dump           # every run length (seed 1), input for plot_runs.py
 *   ./runs time           # CPU time, M = 2^16 and 2^22 records, 5 repetitions
 *   ./runs selftest       # checks both generators against qsort
 */
#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <time.h>

typedef uint64_t key_t64;

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

enum dist { D_RANDOM, D_SORTED, D_REVERSE, D_NEARLY, D_NDIST };
static const char *dist_name[] = {"random", "sorted", "reverse", "nearly-sorted"};

/* nearly-sorted: key = i + uniform[0, 4M), i.e. every key is displaced by less than 4M ranks */
static void gen(key_t64 *a, size_t n, enum dist d, size_t m, uint64_t seed)
{
    rng_state = seed;
    for (size_t i = 0; i < n; i++) {
        switch (d) {
        case D_RANDOM:  a[i] = rng(); break;
        case D_SORTED:  a[i] = i; break;
        case D_REVERSE: a[i] = n - i; break;
        default:        a[i] = i + rng() % (4 * m); break;
        }
    }
}

static uint64_t ncmp; /* key comparisons */

/* ---------- load-sort-store: quicksort each memory load ---------- */
static void isort(key_t64 *a, size_t n)
{
    for (size_t i = 1; i < n; i++) {
        key_t64 x = a[i];
        size_t j = i;
        while (j > 0 && (ncmp++, a[j - 1] > x)) { a[j] = a[j - 1]; j--; }
        a[j] = x;
    }
}

static void qsort_u64(key_t64 *a, size_t n)
{
    while (n > 16) {
        size_t mid = n / 2;
        /* median of three into a[mid] */
        if (ncmp++, a[mid] < a[0]) { key_t64 t = a[mid]; a[mid] = a[0]; a[0] = t; }
        if (ncmp++, a[n - 1] < a[mid]) {
            key_t64 t = a[mid]; a[mid] = a[n - 1]; a[n - 1] = t;
            if (ncmp++, a[mid] < a[0]) { t = a[mid]; a[mid] = a[0]; a[0] = t; }
        }
        key_t64 p = a[mid];
        size_t i = 0, j = n - 1;
        for (;;) { /* Hoare partition */
            while (ncmp++, a[i] < p) i++;
            while (ncmp++, p < a[j]) j--;
            if (i >= j) break;
            key_t64 t = a[i]; a[i] = a[j]; a[j] = t;
            i++; j--;
        }
        /* recurse on the smaller side */
        if (j + 1 < n - j - 1) { qsort_u64(a, j + 1); a += j + 1; n -= j + 1; }
        else { qsort_u64(a + j + 1, n - j - 1); n = j + 1; }
    }
    isort(a, n);
}

/* run lengths are appended to runlen[]; returns number of runs */
static size_t load_sort_store(key_t64 *in, size_t n, size_t m, size_t *runlen, key_t64 *out)
{
    size_t nr = 0;
    for (size_t off = 0; off < n; off += m) {
        size_t len = n - off < m ? n - off : m;
        memcpy(out + off, in + off, len * sizeof *in);
        qsort_u64(out + off, len);
        runlen[nr++] = len;
    }
    return nr;
}

/* ---------- replacement selection ---------- */
typedef struct { uint64_t run; key_t64 key; } hent;

static inline int hless(const hent *a, const hent *b)
{
    ncmp++;
    return a->run < b->run || (a->run == b->run && a->key < b->key);
}

static void sift_down(hent *h, size_t n, size_t i)
{
    hent x = h[i];
    for (;;) {
        size_t c = 2 * i + 1;
        if (c >= n) break;
        if (c + 1 < n && hless(&h[c + 1], &h[c])) c++;
        if (!hless(&h[c], &x)) break;
        h[i] = h[c];
        i = c;
    }
    h[i] = x;
}

static int trace_on;
static void trace_heap(const hent *h, size_t n, uint64_t cur)
{
    /* print heap contents sorted, frozen (next-run) records in brackets */
    hent tmp[16];
    memcpy(tmp, h, n * sizeof *h);
    for (size_t i = 1; i < n; i++)
        for (size_t j = i; j > 0 && (tmp[j].run < tmp[j - 1].run ||
             (tmp[j].run == tmp[j - 1].run && tmp[j].key < tmp[j - 1].key)); j--) {
            hent t = tmp[j]; tmp[j] = tmp[j - 1]; tmp[j - 1] = t;
        }
    for (size_t i = 0; i < n; i++)
        printf(tmp[i].run == cur ? " %llu" : " [%llu]", (unsigned long long)tmp[i].key);
}

static size_t replacement_selection(const key_t64 *in, size_t n, size_t m,
                                    size_t *runlen, key_t64 *out, hent *h)
{
    size_t hn = m < n ? m : n, pos = hn, o = 0, nr = 0;
    for (size_t i = 0; i < hn; i++) h[i] = (hent){0, in[i]};
    for (size_t i = hn / 2; i-- > 0;) sift_down(h, hn, i);
    uint64_t cur = 0;
    size_t len = 0;
    while (hn > 0) {
        hent top = h[0];
        if (top.run != cur) { runlen[nr++] = len; len = 0; cur = top.run; }
        out[o++] = top.key;
        len++;
        if (pos < n) {
            key_t64 x = in[pos++];
            ncmp++; /* x >= last output ? */
            h[0] = (hent){x >= top.key ? cur : cur + 1, x};
            if (trace_on) printf("| %llu | %llu |", (unsigned long long)top.key, (unsigned long long)x);
        } else {
            h[0] = h[--hn];
            if (trace_on) printf("| %llu | - |", (unsigned long long)top.key);
        }
        if (hn > 0) sift_down(h, hn, 0);
        if (trace_on) { trace_heap(h, hn, cur); printf(" | %llu |\n", (unsigned long long)cur); }
    }
    if (len) runlen[nr++] = len;
    return nr;
}

/* ---------- helpers ---------- */
static int cmp_u64(const void *a, const void *b)
{
    key_t64 x = *(const key_t64 *)a, y = *(const key_t64 *)b;
    return (x > y) - (x < y);
}

/* every run must be sorted, and the concatenated runs must be a permutation of the input */
static int check_runs(const key_t64 *in, const key_t64 *out, size_t n, const size_t *runlen, size_t nr)
{
    size_t off = 0;
    for (size_t r = 0; r < nr; r++) {
        for (size_t i = off + 1; i < off + runlen[r]; i++)
            if (out[i - 1] > out[i]) return 0;
        off += runlen[r];
    }
    if (off != n) return 0;
    key_t64 *a = malloc(n * sizeof *a), *b = malloc(n * sizeof *b);
    memcpy(a, in, n * sizeof *a);
    memcpy(b, out, n * sizeof *b);
    qsort(a, n, sizeof *a, cmp_u64);
    qsort(b, n, sizeof *b, cmp_u64);
    int ok = memcmp(a, b, n * sizeof *a) == 0;
    free(a); free(b);
    return ok;
}

static double now(void)
{
    struct timespec ts;
    clock_gettime(CLOCK_PROCESS_CPUTIME_ID, &ts);
    return ts.tv_sec + ts.tv_nsec * 1e-9;
}

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

/* ---------- modes ---------- */
static int selftest(void)
{
    static const size_t ns[] = {0, 1, 2, 3, 10, 1000, 100000};
    static const size_t ms[] = {1, 2, 3, 7, 64, 1000};
    int fails = 0;
    for (size_t a = 0; a < sizeof ns / sizeof *ns; a++)
        for (size_t b = 0; b < sizeof ms / sizeof *ms; b++)
            for (int d = 0; d < D_NDIST; d++)
                for (uint64_t seed = 1; seed <= 3; seed++) {
                    size_t n = ns[a], m = ms[b];
                    key_t64 *in = malloc((n + 1) * sizeof *in), *out = malloc((n + 1) * sizeof *out);
                    size_t *rl = malloc((n + 1) * sizeof *rl);
                    hent *h = malloc((m + 1) * sizeof *h);
                    gen(in, n, d, m, seed);
                    if (d == D_RANDOM && seed == 3) /* many duplicates */
                        for (size_t i = 0; i < n; i++) in[i] %= 5;
                    size_t nr = load_sort_store(in, n, m, rl, out);
                    if (!check_runs(in, out, n, rl, nr)) { printf("FAIL lss n=%zu m=%zu d=%d\n", n, m, d); fails++; }
                    nr = replacement_selection(in, n, m, rl, out, h);
                    if (!check_runs(in, out, n, rl, nr)) { printf("FAIL rs n=%zu m=%zu d=%d\n", n, m, d); fails++; }
                    for (size_t r = 0; r + 1 < nr; r++) /* a non-final run can never be shorter than M */
                        if (rl[r] < m) { printf("FAIL rs short run n=%zu m=%zu\n", n, m); fails++; break; }
                    free(in); free(out); free(rl); free(h);
                }
    printf(fails ? "selftest: %d failures\n" : "selftest: all passed\n", fails);
    return fails != 0;
}

static void trace(void)
{
    static const key_t64 in[] = {50, 20, 80, 30, 90, 10, 70, 40, 60, 5, 15, 25};
    size_t n = sizeof in / sizeof *in, m = 3;
    key_t64 out[16];
    size_t rl[16];
    hent h[4];
    trace_on = 1;
    printf("| out | in | heap after (frozen in []) | run |\n");
    size_t nr = replacement_selection(in, n, m, rl, out, h);
    trace_on = 0;
    printf("runs:");
    for (size_t r = 0; r < nr; r++) printf(" %zu", rl[r]);
    printf("\n");
}

static void lengths(void)
{
    const size_t m = 65536, n = 256 * m;
    key_t64 *in = malloc(n * sizeof *in), *out = malloc(n * sizeof *out);
    size_t *rl = malloc(n * sizeof *rl);
    hent *h = malloc(m * sizeof *h);
    printf("M = %zu, N = %zu (= %zu M)\n", m, n, n / m);
    printf("%-14s %4s | %6s %8s | %6s %9s %9s %9s %9s\n", "input", "seed",
           "LSS", "cmp/rec", "RS", "run1/M", "mid/M", "last/M", "cmp/rec");
    for (int d = 0; d < D_NDIST; d++)
        for (uint64_t seed = 1; seed <= 3; seed++) {
            gen(in, n, d, m, seed);
            ncmp = 0;
            size_t nl = load_sort_store(in, n, m, rl, out);
            double cl = (double)ncmp / n;
            ncmp = 0;
            size_t nr = replacement_selection(in, n, m, rl, out, h);
            double cr = (double)ncmp / n;
            double mid = 0;
            size_t nmid = 0;
            for (size_t r = 1; r + 1 < nr; r++) { mid += rl[r]; nmid++; }
            printf("%-14s %4llu | %6zu %8.2f | %6zu %9.3f ", dist_name[d], (unsigned long long)seed,
                   nl, cl, nr, (double)rl[0] / m);
            if (nmid) printf("%9.3f ", mid / nmid / m); else printf("%9s ", "-");
            if (nr > 1) printf("%9.3f ", (double)rl[nr - 1] / m); else printf("%9s ", "-");
            printf("%9.2f\n", cr);
        }
    free(in); free(out); free(rl); free(h);
}

/* one line per run: input run_index length/M  (replacement selection, seed 1) */
static void dump(void)
{
    const size_t m = 65536, n = 256 * m;
    key_t64 *in = malloc(n * sizeof *in), *out = malloc(n * sizeof *out);
    size_t *rl = malloc(n * sizeof *rl);
    hent *h = malloc(m * sizeof *h);
    for (int d = 0; d < D_NDIST; d++) {
        if (d == D_SORTED) continue;
        gen(in, n, d, m, 1);
        size_t nr = replacement_selection(in, n, m, rl, out, h);
        for (size_t r = 0; r < nr; r++) printf("%s %zu %.4f\n", dist_name[d], r, (double)rl[r] / m);
    }
    free(in); free(out); free(rl); free(h);
}

static void timing(void)
{
    static const size_t ms[] = {1u << 16, 1u << 22};
    for (size_t a = 0; a < 2; a++) {
        size_t m = ms[a], n = 16 * m;
        key_t64 *in = malloc(n * sizeof *in), *out = malloc(n * sizeof *out);
        size_t *rl = malloc(n * sizeof *rl);
        hent *h = malloc(m * sizeof *h);
        gen(in, n, D_RANDOM, m, 1);
        double tl[5], tr[5];
        for (int rep = 0; rep < 5; rep++) {
            double t0 = now();
            load_sort_store(in, n, m, rl, out);
            double t1 = now();
            replacement_selection(in, n, m, rl, out, h);
            double t2 = now();
            tl[rep] = (t1 - t0) / n * 1e9;
            tr[rep] = (t2 - t1) / n * 1e9;
        }
        qsort(tl, 5, sizeof *tl, cmp_d);
        qsort(tr, 5, sizeof *tr, cmp_d);
        printf("M = 2^%d (keys %zu KiB, RS heap %zu KiB), N = 16 M, random: "
               "LSS %.1f ns/rec, RS %.1f ns/rec, RS/LSS = %.2f\n",
               a ? 22 : 16, m * sizeof(key_t64) >> 10, m * sizeof(hent) >> 10,
               tl[2], tr[2], tr[2] / tl[2]);
        free(in); free(out); free(rl); free(h);
    }
}

int main(int argc, char **argv)
{
    const char *mode = argc > 1 ? argv[1] : "lengths";
    if (!strcmp(mode, "selftest")) return selftest();
    if (!strcmp(mode, "trace")) trace();
    else if (!strcmp(mode, "lengths")) lengths();
    else if (!strcmp(mode, "time")) timing();
    else if (!strcmp(mode, "dump")) dump();
    else { fprintf(stderr, "usage: %s trace|lengths|dump|time|selftest\n", argv[0]); return 2; }
    return 0;
}
