
#include <napi.h>
#include <oqs/oqs.h>
#include <vector>
#include <string>

// Helper: convert const uint8_t* to Napi::Buffer
static Napi::Buffer<uint8_t> makeBuffer(Napi::Env env, const uint8_t* data, size_t len) {
  return Napi::Buffer<uint8_t>::Copy(env, data, len);
}

Napi::Value listKEMs(const Napi::CallbackInfo& info) {
  Napi::Env env = info.Env();
  uint32_t count = OQS_KEM_alg_count();
  Napi::Array arr = Napi::Array::New(env, count);
  for (uint32_t i = 0; i < count; ++i) {
    arr.Set(i, Napi::String::New(env, OQS_KEM_alg_identifier(i)));
  }
  return arr;
}

class KemHandle {
public:
  explicit KemHandle(const std::string& alg) {
    kem = OQS_KEM_new(alg.c_str());
    if (!kem) throw std::runtime_error("Unknown KEM algorithm " + alg);
  }
  ~KemHandle() { if (kem) OQS_KEM_free(kem); }
  OQS_KEM* get() { return kem; }
private:
  OQS_KEM* kem;
};

Napi::Value kemKeypair(const Napi::CallbackInfo& info) {
  Napi::Env env = info.Env();
  if (info.Length() < 1 || !info[0].IsString()) {
    Napi::TypeError::New(env, "Expected algorithm name string").ThrowAsJavaScriptException();
    return env.Null();
  }
  std::string alg = info[0].As<Napi::String>().Utf8Value();
  KemHandle kem(alg);
  std::vector<uint8_t> pk(kem.get()->length_public_key);
  std::vector<uint8_t> sk(kem.get()->length_secret_key);
  if (OQS_KEM_keypair(kem.get(), pk.data(), sk.data()) != OQS_SUCCESS)
    Napi::Error::New(env, "OQS_KEM_keypair failed").ThrowAsJavaScriptException();
  Napi::Object result = Napi::Object::New(env);
  result.Set("publicKey", makeBuffer(env, pk.data(), pk.size()));
  result.Set("secretKey", makeBuffer(env, sk.data(), sk.size()));
  return result;
}

Napi::Value encapsulate(const Napi::CallbackInfo& info) {
  Napi::Env env = info.Env();
  if (info.Length() < 2 || !info[0].IsString() || !info[1].IsBuffer()) {
    Napi::TypeError::New(env, "Expected (string alg, Buffer publicKey)").ThrowAsJavaScriptException();
    return env.Null();
  }
  std::string alg = info[0].As<Napi::String>().Utf8Value();
  Napi::Buffer<uint8_t> pkBuf = info[1].As<Napi::Buffer<uint8_t>>();
  KemHandle kem(alg);
  if (pkBuf.Length() != kem.get()->length_public_key) {
    Napi::RangeError::New(env, "publicKey length mismatch").ThrowAsJavaScriptException();
    return env.Null();
  }
  std::vector<uint8_t> ct(kem.get()->length_ciphertext);
  std::vector<uint8_t> ss(kem.get()->length_shared_secret);
  if (OQS_KEM_encaps(kem.get(), ct.data(), ss.data(), pkBuf.Data()) != OQS_SUCCESS)
    Napi::Error::New(env, "OQS_KEM_encaps failed").ThrowAsJavaScriptException();
  Napi::Object result = Napi::Object::New(env);
  result.Set("ciphertext", makeBuffer(env, ct.data(), ct.size()));
  result.Set("sharedSecret", makeBuffer(env, ss.data(), ss.size()));
  return result;
}

Napi::Value decapsulate(const Napi::CallbackInfo& info) {
  Napi::Env env = info.Env();
  if (info.Length() < 3 || !info[0].IsString() || !info[1].IsBuffer() || !info[2].IsBuffer()) {
    Napi::TypeError::New(env, "Expected (string alg, Buffer ciphertext, Buffer secretKey)").ThrowAsJavaScriptException();
    return env.Null();
  }
  std::string alg = info[0].As<Napi::String>().Utf8Value();
  Napi::Buffer<uint8_t> ctBuf = info[1].As<Napi::Buffer<uint8_t>>();
  Napi::Buffer<uint8_t> skBuf = info[2].As<Napi::Buffer<uint8_t>>();
  KemHandle kem(alg);
  if (ctBuf.Length() != kem.get()->length_ciphertext || skBuf.Length() != kem.get()->length_secret_key) {
    Napi::RangeError::New(env, "ciphertext/secretKey length mismatch").ThrowAsJavaScriptException();
    return env.Null();
  }
  std::vector<uint8_t> ss(kem.get()->length_shared_secret);
  if (OQS_KEM_decaps(kem.get(), ss.data(), ctBuf.Data(), skBuf.Data()) != OQS_SUCCESS)
    Napi::Error::New(env, "OQS_KEM_decaps failed").ThrowAsJavaScriptException();
  return makeBuffer(env, ss.data(), ss.size());
}

Napi::Object Init(Napi::Env env, Napi::Object exports) {
  exports.Set("listKEMs", Napi::Function::New(env, listKEMs));
  exports.Set("kemKeypair", Napi::Function::New(env, kemKeypair));
  exports.Set("encapsulate", Napi::Function::New(env, encapsulate));
  exports.Set("decapsulate", Napi::Function::New(env, decapsulate));
  return exports;
}
NODE_API_MODULE(oqs_addon, Init)
