#include "workq.h"
#include <catch2/catch.hpp>
#include <emscripten/atomic.h>
#include <emscripten/wasm_worker.h>

void DeadlineThreadPool::workq_worker_start(int payloadPtr) {
  auto payload = reinterpret_cast<WorkerPayload *>(payloadPtr);

  while (payload->q.doWork())
    ;

  emscripten_atomic_add_u32(&payload->sema, 1);
  emscripten_atomic_notify(&payload->sema, EMSCRIPTEN_NOTIFY_ALL_WAITERS);
}

DeadlineThreadPool::DeadlineThreadPool(int numThreads, DeadlineWorkQ &q)
    : payload(q) {
  workers.reserve(numThreads);
  for (int i = 0; i < numThreads; i++) {
    workers.push_back(emscripten_malloc_wasm_worker(3670016));
    emscripten_wasm_worker_post_function_vi(workers.back(), workq_worker_start,
                                            reinterpret_cast<int>(&payload));
  }
}

DeadlineThreadPool::~DeadlineThreadPool() {
  payload.q.close();

  // wait for all workers
  for (;;) {
    auto v = emscripten_atomic_load_u32(&payload.sema);
    if (v == workers.size()) {
      break;
    }
    emscripten_atomic_wait_u32(&payload.sema, v,
                               ATOMICS_WAIT_DURATION_INFINITE);
  }

  for (auto &worker : workers) {
    emscripten_terminate_wasm_worker(worker);
  }
}

TEST_CASE("DeadlineWorkQ single thread", "[concurrent]") {
  DeadlineWorkQ q;
  int v = 0;
  q.submit([&]() { v = 1; }, 0);
  q.submit([&]() { v = 2; }, 1);
  q.close();
  REQUIRE(q.doWork());
  REQUIRE(v == 1);
  REQUIRE(q.doWork());
  REQUIRE(v == 2);
  REQUIRE(!q.doWork());
}

static constexpr const int NUM_THREADS = 4;

TEST_CASE("DeadlineWorkQ multi thread", "[.][concurrent][noasan]") {
#if defined(__has_feature)
#if __has_feature(address_sanitizer)
  assert(false);
#endif
#endif

  uint32_t v = 0;

  {
    DeadlineWorkQ q{100};
    DeadlineThreadPool tp{NUM_THREADS, q};

    for (int i = 0; i < 1000; i++) {
      q.submit([&] { emscripten_atomic_add_u32(&v, 1); }, 0);
    }
    auto fut = q.submit([] {}, 1);
    q.close();

    fut.wait();
    REQUIRE(fut.isFinished());
  }

  REQUIRE(emscripten_atomic_load_u32(&v) == 1000);
}
