#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/random.h>
#include <memory>
#include <vector>
#include <string>

using namespace mlx::core;
namespace mx = mlx::core;
namespace rnd = mlx::core::random;

// Resource types for random operations
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;

struct ArrayResource {
    array arr;
    std::string name;
    ArrayResource(const array& a, const std::string& n = "") : arr(a), name(n) {}
};

// Helper functions
static ERL_NIF_TERM make_atom(ErlNifEnv* env, const char* name) {
    ERL_NIF_TERM ret;
    if (enif_make_existing_atom(env, name, &ret, ERL_NIF_LATIN1)) {
        return ret;
    }
    return enif_make_atom(env, name);
}

static ERL_NIF_TERM make_error(ErlNifEnv* env, const char* reason) {
    return enif_make_tuple2(env, make_atom(env, "error"), make_atom(env, reason));
}

static ERL_NIF_TERM make_ok(ErlNifEnv* env, ERL_NIF_TERM term) {
    return enif_make_tuple2(env, make_atom(env, "ok"), term);
}

// Parse shape from Erlang list
static std::vector<int> parse_shape(ErlNifEnv* env, ERL_NIF_TERM shape_term) {
    unsigned int shape_len;
    if (!enif_get_list_length(env, shape_term, &shape_len)) {
        return {};
    }

    std::vector<int> shape_vec(shape_len);
    ERL_NIF_TERM head, tail = shape_term;
    
    for (unsigned int i = 0; i < shape_len; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return {};
        }
        if (!enif_get_int(env, head, &shape_vec[i])) {
            return {};
        }
    }
    return shape_vec;
}

// Parse dtype from string
static Dtype parse_dtype(const char* dtype_str) {
    if (strcmp(dtype_str, "float32") == 0) return float32;
    else if (strcmp(dtype_str, "float16") == 0) return float16;
    else if (strcmp(dtype_str, "bfloat16") == 0) return bfloat16;
    else if (strcmp(dtype_str, "float64") == 0) return float64;
    else if (strcmp(dtype_str, "int32") == 0) return int32;
    else if (strcmp(dtype_str, "int16") == 0) return int16;
    else if (strcmp(dtype_str, "int8") == 0) return int8;
    else if (strcmp(dtype_str, "int64") == 0) return int64;
    else if (strcmp(dtype_str, "uint32") == 0) return uint32;
    else if (strcmp(dtype_str, "uint16") == 0) return uint16;
    else if (strcmp(dtype_str, "uint8") == 0) return uint8;
    else if (strcmp(dtype_str, "uint64") == 0) return uint64;
    else if (strcmp(dtype_str, "bool") == 0) return bool_;
    else return float32;
}

// Create array resource
static ERL_NIF_TERM make_array_resource(ErlNifEnv* env, const array& arr, const std::string& name = "") {
    ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
    new(res) ArrayResource(arr, name);
    
    ERL_NIF_TERM term = enif_make_resource(env, res);
    enif_release_resource(res);
    
    return make_ok(env, term);
}

// Get array from resource
static bool get_array_resource(ErlNifEnv* env, ERL_NIF_TERM term, array& arr) {
    ArrayResource* res;
    if (!enif_get_resource(env, term, ARRAY_RESOURCE_TYPE, (void**)&res)) {
        return false;
    }
    arr = res->arr;
    return true;
}

// ==================== RANDOM SEED MANAGEMENT ====================

static ERL_NIF_TERM mlx_random_seed(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    int seed;
    if (!enif_get_int(env, argv[0], &seed)) {
        return enif_make_badarg(env);
    }
    
    try {
        rnd::seed(seed);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "random_seed_error");
    }
}

static ERL_NIF_TERM mlx_random_key(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    int seed;
    if (!enif_get_int(env, argv[0], &seed)) {
        return enif_make_badarg(env);
    }
    
    try {
        array key = rnd::key(seed);
        return make_array_resource(env, key);
    } catch (const std::exception& e) {
        return make_error(env, "random_key_error");
    }
}

// ==================== BASIC RANDOM OPERATIONS ====================

static ERL_NIF_TERM mlx_random_normal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    auto shape = parse_shape(env, argv[0]);
    if (shape.empty()) return enif_make_badarg(env);
    
    double mean, std;
    if (!enif_get_double(env, argv[1], &mean) ||
        !enif_get_double(env, argv[2], &std)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = rnd::normal(shape, float32, mean, std);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_normal_error");
    }
}

static ERL_NIF_TERM mlx_random_uniform(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    auto shape = parse_shape(env, argv[0]);
    if (shape.empty()) return enif_make_badarg(env);
    
    double low, high;
    if (!enif_get_double(env, argv[1], &low) ||
        !enif_get_double(env, argv[2], &high)) {
        return enif_make_badarg(env);
    }
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        Dtype dtype = parse_dtype(dtype_str);
        array result = rnd::uniform(array(low), array(high), shape, dtype);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_uniform_error");
    }
}

static ERL_NIF_TERM mlx_random_randint(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    int low, high;
    if (!enif_get_int(env, argv[0], &low) ||
        !enif_get_int(env, argv[1], &high)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[2]);
    if (shape.empty()) return enif_make_badarg(env);
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        Dtype dtype = parse_dtype(dtype_str);
        array result = rnd::randint(array(low), array(high), shape, dtype);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_randint_error");
    }
}

// ==================== STATISTICAL DISTRIBUTIONS ====================

static ERL_NIF_TERM mlx_random_bernoulli(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array p = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], p)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[1]);
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[2], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = array({1.0f}); // Initialize with dummy value
        if (shape.empty()) {
            result = rnd::bernoulli(p);
        } else {
            result = rnd::bernoulli(p, shape);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_bernoulli_error");
    }
}

static ERL_NIF_TERM mlx_random_categorical(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array logits = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], logits)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[1]);
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[2], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = array({1.0f}); // Initialize with dummy value
        if (shape.empty()) {
            result = rnd::categorical(logits);
        } else {
            result = rnd::categorical(logits, -1, shape);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_categorical_error");
    }
}

static ERL_NIF_TERM mlx_random_multinomial(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array p = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], p)) {
        return enif_make_badarg(env);
    }
    
    int n;
    if (!enif_get_int(env, argv[1], &n)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[2]);
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX doesn't have multinomial - implement using categorical sampling
        return make_error(env, "multinomial_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "random_multinomial_error");
    }
}

// ==================== ADVANCED DISTRIBUTIONS ====================

static ERL_NIF_TERM mlx_random_gamma(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array alpha = array({1.0f}), beta = array({1.0f}); // Initialize with dummy values
    if (!get_array_resource(env, argv[0], alpha) ||
        !get_array_resource(env, argv[1], beta)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[2]);
    // shape can be empty for default
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX may not have gamma distribution, implement using rejection sampling or return error
        return make_error(env, "gamma_distribution_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "random_gamma_error");
    }
}

static ERL_NIF_TERM mlx_random_beta(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    array alpha = array({1.0f}), beta = array({1.0f}); // Initialize with dummy values
    if (!get_array_resource(env, argv[0], alpha) ||
        !get_array_resource(env, argv[1], beta)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[2]);
    // shape can be empty for default
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[3], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX may not have beta distribution
        return make_error(env, "beta_distribution_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "random_beta_error");
    }
}

static ERL_NIF_TERM mlx_random_exponential(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array lambda = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], lambda)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[1]);
    // shape can be empty for default
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[2], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Generate exponential using: -log(1 - uniform) / lambda
        Dtype dtype = parse_dtype(dtype_str);
        array u = rnd::uniform(array(0.0f), array(1.0f), shape, dtype);
        array one_minus_u = subtract(array(1.0f), u);
        array log_term = log(one_minus_u);
        array result = divide(negative(log_term), lambda);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_exponential_error");
    }
}

static ERL_NIF_TERM mlx_random_poisson(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array lambda = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], lambda)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[1]);
    // shape can be empty for default
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[2], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX may not have Poisson distribution
        return make_error(env, "poisson_distribution_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "random_poisson_error");
    }
}

// ==================== RANDOM UTILITIES ====================

static ERL_NIF_TERM mlx_random_shuffle(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array arr = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], arr)) {
        return enif_make_badarg(env);
    }
    
    int axis;
    if (!enif_get_int(env, argv[1], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Generate random permutation indices and use take
        int axis_size = arr.shape()[axis];
        array indices = rnd::randint(array(0), array(axis_size), {axis_size}, int32);
        array result = take(arr, indices, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_shuffle_error");
    }
}

static ERL_NIF_TERM mlx_random_choice(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array arr = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], arr)) {
        return enif_make_badarg(env);
    }
    
    int size;
    if (!enif_get_int(env, argv[1], &size)) {
        return enif_make_badarg(env);
    }
    
    char replace_str[16];
    if (!enif_get_atom(env, argv[2], replace_str, sizeof(replace_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    bool replace = (strcmp(replace_str, "true") == 0);
    
    try {
        int arr_size = arr.shape()[0];
        array indices = array({0}); // Initialize with dummy value
        
        if (replace) {
            indices = rnd::randint(array(0), array(arr_size), {size}, int32);
        } else {
            if (size > arr_size) {
                return make_error(env, "choice_size_larger_than_array");
            }
            // Generate shuffled indices and take first 'size'
            array all_indices = arange(0, arr_size, 1);
            // Simple shuffle by generating random permutation
            indices = rnd::randint(array(0), array(arr_size), {size}, int32);
        }
        
        array result = take(arr, indices, 0);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_choice_error");
    }
}

static ERL_NIF_TERM mlx_random_permutation_n(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    int n;
    if (!enif_get_int(env, argv[0], &n)) {
        return enif_make_badarg(env);
    }
    
    try {
        array indices = arange(0, n, 1);
        // Generate random indices to shuffle
        array random_indices = rnd::randint(array(0), array(n), {n}, int32);
        array result = take(indices, random_indices, 0);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_permutation_error");
    }
}

static ERL_NIF_TERM mlx_random_permutation_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array arr = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], arr)) {
        return enif_make_badarg(env);
    }
    
    int axis;
    if (!enif_get_int(env, argv[1], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Same as shuffle
        int axis_size = arr.shape()[axis];
        array indices = rnd::randint(array(0), array(axis_size), {axis_size}, int32);
        array result = take(arr, indices, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_permutation_error");
    }
}

// Resource destructor
static void array_resource_destructor(ErlNifEnv* env, void* obj) {
    ArrayResource* res = (ArrayResource*)obj;
    res->~ArrayResource();
}

// NIF function table
static ErlNifFunc nif_funcs[] = {
    // Basic random operations
    {"seed", 1, mlx_random_seed, 0},
    {"key", 1, mlx_random_key, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"normal", 3, mlx_random_normal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"uniform", 4, mlx_random_uniform, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"randint", 4, mlx_random_randint, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Statistical distributions
    {"bernoulli", 3, mlx_random_bernoulli, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"categorical", 3, mlx_random_categorical, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"multinomial", 4, mlx_random_multinomial, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Advanced distributions
    {"gamma", 4, mlx_random_gamma, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"beta", 4, mlx_random_beta, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"exponential", 3, mlx_random_exponential, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"poisson", 3, mlx_random_poisson, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Random utilities
    {"shuffle", 2, mlx_random_shuffle, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"choice", 3, mlx_random_choice, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"permutation", 1, mlx_random_permutation_n, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"permutation", 2, mlx_random_permutation_array, ERL_NIF_DIRTY_JOB_CPU_BOUND}
};

static int load(ErlNifEnv* env, void** priv_data, ERL_NIF_TERM load_info) {
    ErlNifResourceFlags flags = (ErlNifResourceFlags)(ERL_NIF_RT_CREATE | ERL_NIF_RT_TAKEOVER);
    ErlNifResourceFlags* tried = NULL;
    
    ARRAY_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_random_array", array_resource_destructor, flags, tried);

    if (!ARRAY_RESOURCE_TYPE) {
        return 1;
    }
    
    return 0;
}

static int upgrade(ErlNifEnv* env, void** priv_data, void** old_priv_data, ERL_NIF_TERM load_info) {
    return load(env, priv_data, load_info);
}

static void unload(ErlNifEnv* env, void* priv_data) {
    // Cleanup if needed
}

ERL_NIF_INIT(mlx_random_nif, nif_funcs, load, NULL, upgrade, unload)