/* Treap (Seidel & Aragon 1996): BST order on key, max-heap order on prio.
 * Two update styles over the same node type:
 *   - rotation-based insert/erase, counting rotations;
 *   - split/merge-based insert/erase.
 * Keys are unique (set semantics). */
#ifndef TREAP_H
#define TREAP_H

#include <stdint.h>
#include <stdlib.h>

typedef struct tnode {
    int64_t key;
    uint64_t prio;
    uint32_t size;
    struct tnode *l, *r;
} tnode;

static inline uint32_t tsize(const tnode *t) { return t ? t->size : 0; }
static inline void pull(tnode *t) { t->size = 1 + tsize(t->l) + tsize(t->r); }

static inline tnode *tnode_new(int64_t key, uint64_t prio)
{
    tnode *x = malloc(sizeof *x);
    if (!x) abort();
    x->key = key;
    x->prio = prio;
    x->size = 1;
    x->l = x->r = NULL;
    return x;
}

static inline void treap_free(tnode *t)
{
    if (!t) return;
    treap_free(t->l);
    treap_free(t->r);
    free(t);
}

/* ---------- rotations ---------- */

static inline tnode *rotate_right(tnode *y)   /* y->l becomes the root */
{
    tnode *x = y->l;
    y->l = x->r;
    x->r = y;
    pull(y);
    pull(x);
    return x;
}

static inline tnode *rotate_left(tnode *x)    /* x->r becomes the root */
{
    tnode *y = x->r;
    x->r = y->l;
    y->l = x;
    pull(x);
    pull(y);
    return y;
}

/* Leaf insertion, then rotate the new node up while its prio beats the parent's.
 * *inserted is set to 0 if the key already exists (x is not linked then). */
static inline tnode *treap_insert_rot(tnode *t, tnode *x, long *rotations, int *inserted)
{
    if (!t) { *inserted = 1; return x; }
    if (x->key == t->key) { *inserted = 0; return t; }
    if (x->key < t->key) {
        t->l = treap_insert_rot(t->l, x, rotations, inserted);
        if (t->l->prio > t->prio) { t = rotate_right(t); ++*rotations; return t; }
    } else {
        t->r = treap_insert_rot(t->r, x, rotations, inserted);
        if (t->r->prio > t->prio) { t = rotate_left(t); ++*rotations; return t; }
    }
    pull(t);
    return t;
}

/* Rotate the target down toward the child with larger prio until it is a leaf. */
static inline tnode *treap_erase_rot(tnode *t, int64_t key, long *rotations, int *erased)
{
    if (!t) { *erased = 0; return NULL; }
    if (key < t->key) {
        t->l = treap_erase_rot(t->l, key, rotations, erased);
    } else if (key > t->key) {
        t->r = treap_erase_rot(t->r, key, rotations, erased);
    } else if (!t->l && !t->r) {
        free(t);
        *erased = 1;
        return NULL;
    } else if (!t->r || (t->l && t->l->prio > t->r->prio)) {
        t = rotate_right(t);
        ++*rotations;
        t->r = treap_erase_rot(t->r, key, rotations, erased);
    } else {
        t = rotate_left(t);
        ++*rotations;
        t->l = treap_erase_rot(t->l, key, rotations, erased);
    }
    pull(t);
    return t;
}

/* ---------- split / merge ---------- */

/* *l gets keys < key, *r gets keys >= key. */
static inline void treap_split(tnode *t, int64_t key, tnode **l, tnode **r)
{
    if (!t) { *l = *r = NULL; return; }
    if (t->key < key) {
        treap_split(t->r, key, &t->r, r);
        *l = t;
    } else {
        treap_split(t->l, key, l, &t->l);
        *r = t;
    }
    pull(t);
}

/* Precondition: every key in l is smaller than every key in r. */
static inline tnode *treap_merge(tnode *l, tnode *r)
{
    if (!l) return r;
    if (!r) return l;
    if (l->prio > r->prio) {
        l->r = treap_merge(l->r, r);
        pull(l);
        return l;
    }
    r->l = treap_merge(l, r->l);
    pull(r);
    return r;
}

/* Descend while the existing nodes have larger prio, then split the rest under x.
 * Precondition: x->key is not in t. */
static inline tnode *treap_insert_sm(tnode *t, tnode *x)
{
    if (!t) return x;
    if (x->prio > t->prio) {
        treap_split(t, x->key, &x->l, &x->r);
        pull(x);
        return x;
    }
    if (x->key < t->key) t->l = treap_insert_sm(t->l, x);
    else                 t->r = treap_insert_sm(t->r, x);
    pull(t);
    return t;
}

/* Replace the target by the merge of its two subtrees. */
static inline tnode *treap_erase_sm(tnode *t, int64_t key, int *erased)
{
    if (!t) { *erased = 0; return NULL; }
    if (key == t->key) {
        tnode *m = treap_merge(t->l, t->r);
        free(t);
        *erased = 1;
        return m;
    }
    if (key < t->key) t->l = treap_erase_sm(t->l, key, erased);
    else              t->r = treap_erase_sm(t->r, key, erased);
    pull(t);
    return t;
}

/* ---------- queries ---------- */

/* Returns the node or NULL; *visited counts nodes on the search path. */
static inline tnode *treap_find(tnode *t, int64_t key, long *visited)
{
    while (t) {
        ++*visited;
        if (key == t->key) return t;
        t = key < t->key ? t->l : t->r;
    }
    return NULL;
}

/* number of keys < key */
static inline uint32_t treap_rank(const tnode *t, int64_t key)
{
    uint32_t r = 0;
    while (t) {
        if (key <= t->key) t = t->l;
        else { r += tsize(t->l) + 1; t = t->r; }
    }
    return r;
}

/* k-th smallest, 1-based; NULL if k is out of range */
static inline const tnode *treap_kth(const tnode *t, uint32_t k)
{
    while (t) {
        uint32_t ls = tsize(t->l);
        if (k <= ls) t = t->l;
        else if (k == ls + 1) return t;
        else { k -= ls + 1; t = t->r; }
    }
    return NULL;
}

/* height in nodes (empty tree = 0) */
static inline int treap_height(const tnode *t)
{
    if (!t) return 0;
    int a = treap_height(t->l), b = treap_height(t->r);
    return 1 + (a > b ? a : b);
}

/* adds depth (in nodes, root = 1) of every node to out[rank-1] */
static inline void treap_depths(const tnode *t, int d, uint32_t base, double *out)
{
    if (!t) return;
    uint32_t me = base + tsize(t->l);
    out[me] += d;
    treap_depths(t->l, d + 1, base, out);
    treap_depths(t->r, d + 1, me + 1, out);
}

#endif
