#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/device.h>
#include <mlx/dtype.h>
#include <mlx/transforms.h>
#include <mlx/random.h>
#include <mlx/linalg.h>
#include <mlx/fft.h>
#include <mlx/io.h>
#include <memory>
#include <vector>
#include <map>
#include <string>
#include <iostream>
#include <functional>
#include <cmath>

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

// Resource types for comprehensive MLX functionality
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;

// Array resource wrapper
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 if (strcmp(dtype_str, "complex64") == 0) return complex64;
    else if (strcmp(dtype_str, "complex128") == 0) return complex128;
    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;
}

// ==================== ARRAY CREATION OPERATIONS ====================

static ERL_NIF_TERM mlx_arange(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double start, stop, step;
    if (!enif_get_double(env, argv[0], &start) ||
        !enif_get_double(env, argv[1], &stop) ||
        !enif_get_double(env, argv[2], &step)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = arange(start, stop, step);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "arange_error");
    }
}

static ERL_NIF_TERM mlx_linspace(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double start, stop;
    int num;
    if (!enif_get_double(env, argv[0], &start) ||
        !enif_get_double(env, argv[1], &stop) ||
        !enif_get_int(env, argv[2], &num)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = linspace(start, stop, num);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "linspace_error");
    }
}

static ERL_NIF_TERM mlx_zeros(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    auto shape = parse_shape(env, argv[0]);
    if (shape.empty()) return enif_make_badarg(env);
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        Dtype dtype = parse_dtype(dtype_str);
        array result = zeros(shape, dtype);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "zeros_error");
    }
}

static ERL_NIF_TERM mlx_ones(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    auto shape = parse_shape(env, argv[0]);
    if (shape.empty()) return enif_make_badarg(env);
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        Dtype dtype = parse_dtype(dtype_str);
        array result = ones(shape, dtype);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "ones_error");
    }
}

static ERL_NIF_TERM mlx_full(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 value;
    if (!enif_get_double(env, argv[1], &value)) {
        return enif_make_badarg(env);
    }
    
    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 {
        Dtype dtype = parse_dtype(dtype_str);
        array result = full(shape, array(value), dtype);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "full_error");
    }
}

static ERL_NIF_TERM mlx_eye(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    int n;
    if (!enif_get_int(env, argv[0], &n)) {
        return enif_make_badarg(env);
    }
    
    char dtype_str[32];
    if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        Dtype dtype = parse_dtype(dtype_str);
        array result = eye(n, dtype);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "eye_error");
    }
}

// ==================== ARITHMETIC OPERATIONS ====================

static ERL_NIF_TERM mlx_add(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = add(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "add_error");
    }
}

static ERL_NIF_TERM mlx_subtract(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = subtract(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "subtract_error");
    }
}

static ERL_NIF_TERM mlx_multiply(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = multiply(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "multiply_error");
    }
}

static ERL_NIF_TERM mlx_divide(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = divide(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "divide_error");
    }
}

static ERL_NIF_TERM mlx_power(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = power(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "power_error");
    }
}

static ERL_NIF_TERM mlx_negative(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = negative(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "negative_error");
    }
}

// ==================== TRIGONOMETRIC FUNCTIONS ====================

static ERL_NIF_TERM mlx_sin(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = sin(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sin_error");
    }
}

static ERL_NIF_TERM mlx_cos(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = cos(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "cos_error");
    }
}

static ERL_NIF_TERM mlx_tan(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = tan(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "tan_error");
    }
}

static ERL_NIF_TERM mlx_arcsin(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = arcsin(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "arcsin_error");
    }
}

static ERL_NIF_TERM mlx_arccos(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = arccos(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "arccos_error");
    }
}

static ERL_NIF_TERM mlx_arctan(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = arctan(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "arctan_error");
    }
}

// ==================== HYPERBOLIC FUNCTIONS ====================

static ERL_NIF_TERM mlx_sinh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = sinh(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sinh_error");
    }
}

static ERL_NIF_TERM mlx_cosh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = cosh(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "cosh_error");
    }
}

static ERL_NIF_TERM mlx_tanh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = tanh(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "tanh_error");
    }
}

// ==================== EXPONENTIAL AND LOGARITHMIC ====================

static ERL_NIF_TERM mlx_exp(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = exp(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "exp_error");
    }
}

static ERL_NIF_TERM mlx_log(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = log(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "log_error");
    }
}

static ERL_NIF_TERM mlx_log2(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = log2(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "log2_error");
    }
}

static ERL_NIF_TERM mlx_log10(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = log10(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "log10_error");
    }
}

static ERL_NIF_TERM mlx_sqrt(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = sqrt(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sqrt_error");
    }
}

static ERL_NIF_TERM mlx_square(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = square(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "square_error");
    }
}

// ==================== COMPARISON OPERATIONS ====================

static ERL_NIF_TERM mlx_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = equal(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "equal_error");
    }
}

static ERL_NIF_TERM mlx_not_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = not_equal(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "not_equal_error");
    }
}

static ERL_NIF_TERM mlx_greater(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = greater(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "greater_error");
    }
}

static ERL_NIF_TERM mlx_greater_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = greater_equal(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "greater_equal_error");
    }
}

static ERL_NIF_TERM mlx_less(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = less(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "less_error");
    }
}

static ERL_NIF_TERM mlx_less_equal(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = less_equal(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "less_equal_error");
    }
}

// ==================== LOGICAL OPERATIONS ====================

static ERL_NIF_TERM mlx_logical_and(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = logical_and(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "logical_and_error");
    }
}

static ERL_NIF_TERM mlx_logical_or(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = logical_or(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "logical_or_error");
    }
}

static ERL_NIF_TERM mlx_logical_not(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = logical_not(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "logical_not_error");
    }
}

// ==================== REDUCTION OPERATIONS ====================

static ERL_NIF_TERM mlx_sum(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 3) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = sum(a);
        } else if (argc == 2) {
            // Sum along specific axes
            auto axes = parse_shape(env, argv[1]);
            result = sum(a, axes);
        } else {
            // Sum with keepdims
            auto axes = parse_shape(env, argv[1]);
            int keepdims;
            if (!enif_get_int(env, argv[2], &keepdims)) {
                return enif_make_badarg(env);
            }
            result = sum(a, axes, keepdims != 0);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sum_error");
    }
}

static ERL_NIF_TERM mlx_mean(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 3) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = mean(a);
        } else if (argc == 2) {
            auto axes = parse_shape(env, argv[1]);
            result = mean(a, axes);
        } else {
            auto axes = parse_shape(env, argv[1]);
            int keepdims;
            if (!enif_get_int(env, argv[2], &keepdims)) {
                return enif_make_badarg(env);
            }
            result = mean(a, axes, keepdims != 0);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "mean_error");
    }
}

static ERL_NIF_TERM mlx_max(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 3) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = max(a);
        } else if (argc == 2) {
            auto axes = parse_shape(env, argv[1]);
            result = max(a, axes);
        } else {
            auto axes = parse_shape(env, argv[1]);
            int keepdims;
            if (!enif_get_int(env, argv[2], &keepdims)) {
                return enif_make_badarg(env);
            }
            result = max(a, axes, keepdims != 0);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "max_error");
    }
}

static ERL_NIF_TERM mlx_min(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 3) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = min(a);
        } else if (argc == 2) {
            auto axes = parse_shape(env, argv[1]);
            result = min(a, axes);
        } else {
            auto axes = parse_shape(env, argv[1]);
            int keepdims;
            if (!enif_get_int(env, argv[2], &keepdims)) {
                return enif_make_badarg(env);
            }
            result = min(a, axes, keepdims != 0);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "min_error");
    }
}

static ERL_NIF_TERM mlx_var(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 4) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = var(a);
        } else {
            auto axes = parse_shape(env, argv[1]);
            int keepdims = 0, ddof = 0;
            if (argc > 2 && !enif_get_int(env, argv[2], &keepdims)) {
                return enif_make_badarg(env);
            }
            if (argc > 3 && !enif_get_int(env, argv[3], &ddof)) {
                return enif_make_badarg(env);
            }
            result = var(a, axes, keepdims != 0, ddof);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "var_error");
    }
}

static ERL_NIF_TERM mlx_std(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 4) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = std(a);
        } else {
            auto axes = parse_shape(env, argv[1]);
            int keepdims = 0, ddof = 0;
            if (argc > 2 && !enif_get_int(env, argv[2], &keepdims)) {
                return enif_make_badarg(env);
            }
            if (argc > 3 && !enif_get_int(env, argv[3], &ddof)) {
                return enif_make_badarg(env);
            }
            result = std(a, axes, keepdims != 0, ddof);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "std_error");
    }
}

// ==================== SHAPE MANIPULATION ====================

static ERL_NIF_TERM mlx_reshape(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[1]);
    if (shape.empty()) return enif_make_badarg(env);
    
    try {
        array result = reshape(a, shape);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "reshape_error");
    }
}

static ERL_NIF_TERM mlx_transpose(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = transpose(a);
        } else {
            auto axes = parse_shape(env, argv[1]);
            result = transpose(a, axes);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "transpose_error");
    }
}

static ERL_NIF_TERM mlx_squeeze(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = squeeze(a);
        } else {
            auto axes = parse_shape(env, argv[1]);
            result = squeeze(a, axes);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "squeeze_error");
    }
}

static ERL_NIF_TERM mlx_expand_dims(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    auto axes = parse_shape(env, argv[1]);
    if (axes.empty()) return enif_make_badarg(env);
    
    try {
        array result = expand_dims(a, axes);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "expand_dims_error");
    }
}

// ==================== ARRAY CONCATENATION AND STACKING ====================

static ERL_NIF_TERM mlx_concatenate(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    // Parse array list
    unsigned int array_count;
    if (!enif_get_list_length(env, argv[0], &array_count)) {
        return enif_make_badarg(env);
    }
    
    std::vector<array> arrays;
    arrays.reserve(array_count);
    
    ERL_NIF_TERM head, tail = argv[0];
    for (unsigned int i = 0; i < array_count; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return enif_make_badarg(env);
        }
        array arr;
        if (!get_array_resource(env, head, arr)) {
            return enif_make_badarg(env);
        }
        arrays.push_back(arr);
    }
    
    int axis;
    if (!enif_get_int(env, argv[1], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = concatenate(arrays, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "concatenate_error");
    }
}

static ERL_NIF_TERM mlx_stack(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    // Parse array list
    unsigned int array_count;
    if (!enif_get_list_length(env, argv[0], &array_count)) {
        return enif_make_badarg(env);
    }
    
    std::vector<array> arrays;
    arrays.reserve(array_count);
    
    ERL_NIF_TERM head, tail = argv[0];
    for (unsigned int i = 0; i < array_count; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return enif_make_badarg(env);
        }
        array arr;
        if (!get_array_resource(env, head, arr)) {
            return enif_make_badarg(env);
        }
        arrays.push_back(arr);
    }
    
    int axis;
    if (!enif_get_int(env, argv[1], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = stack(arrays, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "stack_error");
    }
}

// ==================== LINEAR ALGEBRA ====================

static ERL_NIF_TERM mlx_matmul(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array a, b;
    if (!get_array_resource(env, argv[0], a) || !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = matmul(a, b);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "matmul_error");
    }
}

// ==================== SORTING AND SEARCHING ====================

static ERL_NIF_TERM mlx_sort(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = sort(a);
        } else {
            int axis;
            if (!enif_get_int(env, argv[1], &axis)) {
                return enif_make_badarg(env);
            }
            result = sort(a, axis);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sort_error");
    }
}

static ERL_NIF_TERM mlx_argsort(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result;
        if (argc == 1) {
            result = argsort(a);
        } else {
            int axis;
            if (!enif_get_int(env, argv[1], &axis)) {
                return enif_make_badarg(env);
            }
            result = argsort(a, axis);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "argsort_error");
    }
}

// ==================== UTILITY FUNCTIONS ====================

static ERL_NIF_TERM mlx_where(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    array condition, x, y;
    if (!get_array_resource(env, argv[0], condition) ||
        !get_array_resource(env, argv[1], x) ||
        !get_array_resource(env, argv[2], y)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = where(condition, x, y);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "where_error");
    }
}

static ERL_NIF_TERM mlx_abs(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = abs(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "abs_error");
    }
}

static ERL_NIF_TERM mlx_sign(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = sign(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "sign_error");
    }
}

// ==================== ARRAY INFO FUNCTIONS ====================

static ERL_NIF_TERM mlx_shape(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        const std::vector<int>& shape = a.shape();
        ERL_NIF_TERM shape_list = enif_make_list(env, 0);
        
        for (int i = shape.size() - 1; i >= 0; i--) {
            shape_list = enif_make_list_cell(env, enif_make_int(env, shape[i]), shape_list);
        }
        
        return make_ok(env, shape_list);
    } catch (const std::exception& e) {
        return make_error(env, "shape_error");
    }
}

static ERL_NIF_TERM mlx_size(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        size_t size = a.size();
        return make_ok(env, enif_make_ulong(env, size));
    } catch (const std::exception& e) {
        return make_error(env, "size_error");
    }
}

static ERL_NIF_TERM mlx_ndim(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        int ndim = a.ndim();
        return make_ok(env, enif_make_int(env, ndim));
    } catch (const std::exception& e) {
        return make_error(env, "ndim_error");
    }
}

static ERL_NIF_TERM mlx_dtype_str(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        Dtype dtype = a.dtype();
        const char* dtype_name;
        
        if (dtype == float32) dtype_name = "float32";
        else if (dtype == float16) dtype_name = "float16";
        else if (dtype == bfloat16) dtype_name = "bfloat16";
        else if (dtype == float64) dtype_name = "float64";
        else if (dtype == int32) dtype_name = "int32";
        else if (dtype == int16) dtype_name = "int16";
        else if (dtype == int8) dtype_name = "int8";
        else if (dtype == int64) dtype_name = "int64";
        else if (dtype == uint32) dtype_name = "uint32";
        else if (dtype == uint16) dtype_name = "uint16";
        else if (dtype == uint8) dtype_name = "uint8";
        else if (dtype == uint64) dtype_name = "uint64";
        else if (dtype == bool_) dtype_name = "bool";
        else if (dtype == complex64) dtype_name = "complex64";
        else if (dtype == complex128) dtype_name = "complex128";
        else dtype_name = "unknown";
        
        return make_ok(env, make_atom(env, dtype_name));
    } catch (const std::exception& e) {
        return make_error(env, "dtype_error");
    }
}

// ==================== EVALUATION ====================

static ERL_NIF_TERM mlx_eval(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a;
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        eval(a);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "eval_error");
    }
}

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

// Complete MLX NIF function table
static ErlNifFunc nif_funcs[] = {
    // Array creation
    {"arange", 3, mlx_arange, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"linspace", 3, mlx_linspace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"zeros", 2, mlx_zeros, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"ones", 2, mlx_ones, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"full", 3, mlx_full, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"eye", 2, mlx_eye, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Arithmetic operations
    {"add", 2, mlx_add, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"subtract", 2, mlx_subtract, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"multiply", 2, mlx_multiply, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"divide", 2, mlx_divide, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"power", 2, mlx_power, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"negative", 1, mlx_negative, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Trigonometric functions
    {"sin", 1, mlx_sin, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"cos", 1, mlx_cos, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"tan", 1, mlx_tan, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"arcsin", 1, mlx_arcsin, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"arccos", 1, mlx_arccos, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"arctan", 1, mlx_arctan, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Hyperbolic functions
    {"sinh", 1, mlx_sinh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"cosh", 1, mlx_cosh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"tanh", 1, mlx_tanh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Exponential and logarithmic
    {"exp", 1, mlx_exp, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"log", 1, mlx_log, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"log2", 1, mlx_log2, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"log10", 1, mlx_log10, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"sqrt", 1, mlx_sqrt, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"square", 1, mlx_square, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Comparison operations
    {"equal", 2, mlx_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"not_equal", 2, mlx_not_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"greater", 2, mlx_greater, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"greater_equal", 2, mlx_greater_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"less", 2, mlx_less, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"less_equal", 2, mlx_less_equal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Logical operations
    {"logical_and", 2, mlx_logical_and, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"logical_or", 2, mlx_logical_or, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"logical_not", 1, mlx_logical_not, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Reduction operations
    {"sum", 1, mlx_sum, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"sum", 2, mlx_sum, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"sum", 3, mlx_sum, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"mean", 1, mlx_mean, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"mean", 2, mlx_mean, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"mean", 3, mlx_mean, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"max", 1, mlx_max, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"max", 2, mlx_max, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"max", 3, mlx_max, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"min", 1, mlx_min, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"min", 2, mlx_min, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"min", 3, mlx_min, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"var", 1, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"var", 2, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"var", 3, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"var", 4, mlx_var, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"std", 1, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"std", 2, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"std", 3, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"std", 4, mlx_std, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Shape manipulation
    {"reshape", 2, mlx_reshape, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"transpose", 1, mlx_transpose, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"transpose", 2, mlx_transpose, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"squeeze", 1, mlx_squeeze, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"squeeze", 2, mlx_squeeze, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"expand_dims", 2, mlx_expand_dims, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Array concatenation and stacking
    {"concatenate", 2, mlx_concatenate, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"stack", 2, mlx_stack, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Linear algebra
    {"matmul", 2, mlx_matmul, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Sorting and searching
    {"sort", 1, mlx_sort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"sort", 2, mlx_sort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"argsort", 1, mlx_argsort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"argsort", 2, mlx_argsort, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Utility functions
    {"where", 3, mlx_where, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"abs", 1, mlx_abs, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"sign", 1, mlx_sign, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Array info
    {"shape", 1, mlx_shape, 0},
    {"size", 1, mlx_size, 0},
    {"ndim", 1, mlx_ndim, 0},
    {"dtype_str", 1, mlx_dtype_str, 0},
    
    // Evaluation
    {"eval", 1, mlx_eval, 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_complete_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_complete_nif, nif_funcs, load, NULL, upgrade, unload)