/*
 * Lazy (optimistic, lock-based) skip-list set after Herlihy, Lev,
 * Luchangco, Shavit, "A Simple Optimistic Skiplist Algorithm" (SIROCCO
 * 2007). Searches take no locks; updates lock predecessors bottom-up,
 * validate, then link or unlink. A key is in the set iff its node is
 * reachable, fully_linked and not marked.
 */
#include <pthread.h>
#include "common.h"

typedef struct lz_node {
    int64_t key;
    int top;
    pthread_mutex_t lock;
    atomic_bool marked;
    atomic_bool fully_linked;
    _Atomic(struct lz_node *) next[];
} lz_node_t;

typedef struct {
    lz_node_t *head, *tail;
} lz_t;

static inline lz_node_t *ld(_Atomic(lz_node_t *) *a)
{
    return atomic_load_explicit(a, memory_order_acquire);
}

static inline void st(_Atomic(lz_node_t *) *a, lz_node_t *v)
{
    atomic_store_explicit(a, v, memory_order_release);
}

static lz_node_t *node_new(ctx_t *c, int64_t key, int top)
{
    size_t sz = sizeof(lz_node_t) + top * sizeof(_Atomic(lz_node_t *));
    lz_node_t *n = c ? ctx_alloc(c, sz) : calloc(1, sz);
    n->key = key;
    n->top = top;
    pthread_mutex_init(&n->lock, NULL);
    return n;
}

static void *lz_create(void)
{
    lz_t *s = malloc(sizeof(*s));
    s->head = node_new(NULL, KEY_MIN, MAX_LEVEL);
    s->tail = node_new(NULL, KEY_MAX, MAX_LEVEL);
    for (int i = 0; i < MAX_LEVEL; i++)
        atomic_init(&s->head->next[i], s->tail);
    atomic_init(&s->head->fully_linked, 1);
    atomic_init(&s->tail->fully_linked, 1);
    return s;
}

static void lz_destroy(void *p)
{
    lz_t *s = p;
    free(s->head);
    free(s->tail);
    free(s);
}

/* Returns the highest level at which key was found, or -1. */
static int find(lz_t *s, int64_t key, lz_node_t **preds, lz_node_t **succs)
{
    int found = -1;
    lz_node_t *pred = s->head;
    for (int lv = MAX_LEVEL - 1; lv >= 0; lv--) {
        lz_node_t *curr = ld(&pred->next[lv]);
        while (curr->key < key) {
            pred = curr;
            curr = ld(&pred->next[lv]);
        }
        if (found == -1 && curr->key == key)
            found = lv;
        preds[lv] = pred;
        succs[lv] = curr;
    }
    return found;
}

/* Unlock each distinct predecessor once (the same node can be pred at many levels). */
static void unlock_preds(lz_node_t **preds, int highest)
{
    lz_node_t *prev = NULL;
    for (int lv = 0; lv <= highest; lv++) {
        if (preds[lv] != prev) {
            pthread_mutex_unlock(&preds[lv]->lock);
            prev = preds[lv];
        }
    }
}

static int lz_insert(void *p, ctx_t *c, int64_t key)
{
    lz_t *s = p;
    int top = random_level(c);
    lz_node_t *preds[MAX_LEVEL], *succs[MAX_LEVEL];
    for (;;) {
        int f = find(s, key, preds, succs);
        if (f != -1) {
            lz_node_t *hit = succs[f];
            if (!atomic_load(&hit->marked)) {
                while (!atomic_load(&hit->fully_linked))
                    ;                 /* becomes a member once fully linked */
                return 0;
            }
            c->retry++;
            continue;                 /* being removed: retry */
        }
        int highest = -1, valid = 1;
        lz_node_t *prev = NULL;
        for (int lv = 0; valid && lv < top; lv++) {
            lz_node_t *pred = preds[lv], *succ = succs[lv];
            if (pred != prev) {
                pthread_mutex_lock(&pred->lock);
                prev = pred;
            }
            highest = lv;
            valid = !atomic_load(&pred->marked) && !atomic_load(&succ->marked) &&
                    ld(&pred->next[lv]) == succ;
        }
        if (!valid) {
            unlock_preds(preds, highest);
            c->retry++;
            continue;
        }
        lz_node_t *n = node_new(c, key, top);
        for (int lv = 0; lv < top; lv++)
            atomic_store_explicit(&n->next[lv], succs[lv], memory_order_relaxed);
        for (int lv = 0; lv < top; lv++)
            st(&preds[lv]->next[lv], n);
        atomic_store(&n->fully_linked, 1);    /* linearization point */
        unlock_preds(preds, highest);
        return 1;
    }
}

static int ok_to_delete(lz_node_t *n, int found)
{
    return atomic_load(&n->fully_linked) && n->top - 1 == found &&
           !atomic_load(&n->marked);
}

static int lz_remove(void *p, ctx_t *c, int64_t key)
{
    lz_t *s = p;
    lz_node_t *victim = NULL;
    int is_marked = 0, top = -1;
    lz_node_t *preds[MAX_LEVEL], *succs[MAX_LEVEL];
    for (;;) {
        int f = find(s, key, preds, succs);
        if (!is_marked && !(f != -1 && ok_to_delete(succs[f], f)))
            return 0;
        if (!is_marked) {
            victim = succs[f];
            top = victim->top;
            pthread_mutex_lock(&victim->lock);
            if (atomic_load(&victim->marked)) {
                pthread_mutex_unlock(&victim->lock);
                return 0;
            }
            atomic_store(&victim->marked, 1); /* linearization point */
            is_marked = 1;
        }
        int highest = -1, valid = 1;
        lz_node_t *prev = NULL;
        for (int lv = 0; valid && lv < top; lv++) {
            lz_node_t *pred = preds[lv];
            if (pred != prev) {
                pthread_mutex_lock(&pred->lock);
                prev = pred;
            }
            highest = lv;
            valid = !atomic_load(&pred->marked) && ld(&pred->next[lv]) == victim;
        }
        if (!valid) {
            unlock_preds(preds, highest);
            c->retry++;
            continue;
        }
        for (int lv = top - 1; lv >= 0; lv--)
            st(&preds[lv]->next[lv], ld(&victim->next[lv]));
        pthread_mutex_unlock(&victim->lock);
        unlock_preds(preds, highest);
        return 1;
    }
}

static int lz_contains(void *p, ctx_t *c, int64_t key)
{
    (void)c;
    lz_node_t *preds[MAX_LEVEL], *succs[MAX_LEVEL];
    int f = find(p, key, preds, succs);
    return f != -1 && atomic_load(&succs[f]->fully_linked) &&
           !atomic_load(&succs[f]->marked);
}

static long lz_check(void *p, int64_t *out, long cap)
{
    lz_t *s = p;
    long n = 0;
    for (int lv = MAX_LEVEL - 1; lv >= 0; lv--) {
        int64_t prev = KEY_MIN;
        lz_node_t *below = lv > 0 ? ld(&s->head->next[lv - 1]) : NULL;
        for (lz_node_t *x = ld(&s->head->next[lv]); x != s->tail; x = ld(&x->next[lv])) {
            if (atomic_load(&x->marked) || !atomic_load(&x->fully_linked) ||
                x->key <= prev || lv >= x->top)
                return -1;
            if (lv > 0) {
                while (below != s->tail && below->key < x->key)
                    below = ld(&below->next[lv - 1]);
                if (below != x)
                    return -1;
            } else {
                if (n < cap)
                    out[n] = x->key;
                n++;
            }
            prev = x->key;
        }
    }
    return n;
}

const set_ops_t lazy_ops = {
    "lazy", lz_create, lz_destroy, lz_insert, lz_remove, lz_contains, lz_check,
};
