// (c) Yura Zhivaga <yzhivaga@gmail.com>
// 2018

#include <vector>
#include <string>
#include <boost/multiprecision/cpp_int.hpp>

#include "nifpp.h"
#include "srp.hh"
#include "sha/sha1.hh"
#include "sha/sha512.hh"

using Int = boost::multiprecision::number<boost::multiprecision::cpp_int_backend<>, boost::multiprecision::et_off>;

SRP *srp;

extern "C" {

static ERL_NIF_TERM erl_compute_s(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    return nifpp::make(env, srp->s());    
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_compute_v(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string x = nifpp::get<std::string>(env, argv[0]);
    std::string v = srp->v(x);
    return nifpp::make(env, v);
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_compute_u(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string A = nifpp::get<std::string>(env, argv[0]);
    std::string B = nifpp::get<std::string>(env, argv[1]);
    std::string u = srp->u(A, B);
    return nifpp::make(env, u);    
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_compute_B(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string v = nifpp::get<std::string>(env, argv[0]);
    std::string b = srp->b();
    std::string B = srp->B(v, b); 
    return nifpp::make(env, std::make_tuple(B, b));
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_compute_K(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string A = nifpp::get<std::string>(env, argv[0]);
    std::string v = nifpp::get<std::string>(env, argv[1]);
    std::string u = nifpp::get<std::string>(env, argv[2]);
    std::string b = nifpp::get<std::string>(env, argv[3]);
    std::string server_k = srp->H({srp->server_s(A, v, u, b)});
    return nifpp::make(env, server_k);
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_compute_M(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string I = nifpp::get<std::string>(env, argv[0]);
    std::string s = nifpp::get<std::string>(env, argv[1]);
    std::string A = nifpp::get<std::string>(env, argv[2]);
    std::string B = nifpp::get<std::string>(env, argv[3]);
    std::string K = nifpp::get<std::string>(env, argv[4]);
    std::string server_M = srp->M(I, s, A, B, K);
    return nifpp::make(env, server_M);
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_compute_R(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string A = nifpp::get<std::string>(env, argv[0]);
    std::string M = nifpp::get<std::string>(env, argv[1]);
    std::string K = nifpp::get<std::string>(env, argv[2]);
    std::string server_R = srp->R(A, M, K);
    return nifpp::make(env, server_R);
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_is_zero_mod_N(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string number = nifpp::get<std::string>(env, argv[0]);
    return nifpp::make(env, srp->is_zero_mod_n(number));
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_sha1(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string input = nifpp::get<std::string>(env, argv[0]);
    std::string output = sw::sha1::calculate(input);

    return nifpp::make(env, output);
  } catch(...) {}
  return enif_make_badarg(env);
}

static ERL_NIF_TERM erl_sha512(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
  try {
    std::string input = nifpp::get<std::string>(env, argv[0]);
    std::string output = sw::sha512::calculate(input);

    return nifpp::make(env, output);
  } catch(...) {}
  return enif_make_badarg(env);
}

static ErlNifFunc nif_funcs[] = {
  {"sha1", 1, erl_sha1},
  {"sha512", 1, erl_sha512},
  {"compute_B", 1, erl_compute_B},
  {"compute_u", 2, erl_compute_u},
  {"compute_K", 4, erl_compute_K},
  {"compute_M", 5, erl_compute_M},
  {"compute_R", 3, erl_compute_R},
  {"compute_s", 0, erl_compute_s},
  {"compute_v", 1, erl_compute_v},
  {"is_zero_mod_N", 1, erl_is_zero_mod_N}
};

static ErlNifResourceType *MEM_RESOURCE;

void reload_srp(ErlNifEnv* env, ERL_NIF_TERM load_info) {
  std::vector<std::string> argv;
  nifpp::get_throws(env, load_info, argv);

  Int N = Int(argv[0]);
  Int g = Int(argv[1]);
  Int k = Int(argv[2]);
  srp = new SRP(N, g, k);
}

static int load(ErlNifEnv* env, void** priv, ERL_NIF_TERM load_info) {
  reload_srp(env, load_info);
  return 0;
}

static int upgrade(ErlNifEnv* env, void** priv, void** old_priv, ERL_NIF_TERM load_info) {
  reload_srp(env, load_info);
  return 0;
}

static int reload(ErlNifEnv* env, void** priv, ERL_NIF_TERM load_info) {
  reload_srp(env, load_info);
  return 0;
}

static void unload(ErlNifEnv* env, void* priv) {
  return;
}

ERL_NIF_INIT(Elixir.CryptxC, nif_funcs, &load, &reload, &upgrade, &unload)

} // extern C
