#pragma once
#include "grip.h"
#include "machine.h"
#include <concepts>
#include <cstdint>
#include <emscripten/atomic.h>

static_assert(sizeof(int *) == sizeof(uint32_t));

template <typename Node>
  requires requires(Node n) {
    { n.next } -> std::same_as<Node *&>;
  }
class alignas(CACHE_LINE_SIZE) IntrusiveBag {
  Grip grip;
  alignas(CACHE_LINE_SIZE) uint64_t head;

  Node *untagPointer(uint64_t tagged) {
    return reinterpret_cast<Node *>(tagged & 0xFFFFFFFFULL);
  }

  uint64_t incrementRetagPointer(uint64_t oldTagged, Node *newPointer) {
    return ((oldTagged + 0x100000000ULL) & 0xFFFFFFFF00000000ULL) |
           reinterpret_cast<uint32_t>(newPointer);
  }

public:
  IntrusiveBag(const IntrusiveBag &) = delete;
  IntrusiveBag &operator=(const IntrusiveBag &) = delete;
  IntrusiveBag(IntrusiveBag &&) = delete;
  IntrusiveBag &operator=(IntrusiveBag &&) = delete;

  IntrusiveBag(std::function<bool()> &&realtimePred)
      : grip(std::move(realtimePred)), head(0) {}

  void add(Node *n) {
    grip.with([this, n] {
      auto currentHead = emscripten_atomic_load_u64(&head);
      emscripten_atomic_store_u32(
          &(n->next), reinterpret_cast<uint32_t>(untagPointer(currentHead)));
      return emscripten_atomic_cas_u64(&head, currentHead,
                                       incrementRetagPointer(currentHead, n)) ==
             currentHead;
    });
  }

  Node *remove() {
    Node *node = nullptr;
    grip.with([this, &node] {
      auto currentHead = emscripten_atomic_load_u64(&head);
      node = untagPointer(currentHead);
      if (node == nullptr) {
        return true;
      }
      auto nextNode =
          reinterpret_cast<Node *>(emscripten_atomic_load_u32(&(node->next)));
      if (emscripten_atomic_cas_u64(
              &head, currentHead,
              incrementRetagPointer(currentHead, nextNode)) != currentHead) {
        return false;
      }
      node->next = nullptr;
      return true;
    });
    return node;
  }

  Node *removeAll() {
    Node *node = nullptr;
    grip.with([this, &node] {
      auto currentHead = emscripten_atomic_load_u64(&head);
      node = untagPointer(currentHead);
      if (node == nullptr) {
        return true;
      }
      return emscripten_atomic_cas_u64(
                 &head, currentHead,
                 incrementRetagPointer(currentHead, nullptr)) == currentHead;
    });
    return node;
  }

  void notifyHead() {
    emscripten_atomic_notify(&head, EMSCRIPTEN_NOTIFY_ALL_WAITERS);
  }

  // requires notifyHead to be called whenever an element is added
  void waitNonempty() {
    for (;;) {
      auto currentHead = emscripten_atomic_load_u64(&head);
      if (untagPointer(currentHead) != nullptr) {
        return;
      }
      emscripten_atomic_wait_u64(&head, currentHead,
                                 ATOMICS_WAIT_DURATION_INFINITE);
    }
  }
};