#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;
    // complex128 is not supported in MLX, map to complex64
    else if (strcmp(dtype_str, "complex128") == 0) return complex64;
    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 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;
}

// Get array from resource (by value)
static array get_array_resource_val(ErlNifEnv* env, ERL_NIF_TERM term, bool& success) {
    ArrayResource* res;
    if (!enif_get_resource(env, term, ARRAY_RESOURCE_TYPE, (void**)&res)) {
        success = false;
        return array({0.0f}); // Return a dummy array
    }
    success = true;
    return res->arr;
}

// ==================== 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");
    }
}

// Helper function to parse nested lists into a flat vector and determine shape
static bool parse_nested_list(ErlNifEnv* env, ERL_NIF_TERM term, std::vector<float>& data, std::vector<int>& shape, int depth = 0) {
    unsigned int list_len;
    if (!enif_get_list_length(env, term, &list_len)) {
        // Not a list, must be a number
        double val;
        int int_val;
        if (enif_get_double(env, term, &val)) {
            data.push_back(static_cast<float>(val));
            return true;
        } else if (enif_get_int(env, term, &int_val)) {
            data.push_back(static_cast<float>(int_val));
            return true;
        }
        return false;
    }
    
    if (depth >= shape.size()) {
        shape.push_back(list_len);
    } else if (shape[depth] != static_cast<int>(list_len)) {
        // Shape mismatch
        return false;
    }
    
    ERL_NIF_TERM head, tail = term;
    for (unsigned int i = 0; i < list_len; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return false;
        }
        if (!parse_nested_list(env, head, data, shape, depth + 1)) {
            return false;
        }
    }
    
    return true;
}

static ERL_NIF_TERM mlx_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    // Parse the data
    std::vector<float> data;
    std::vector<int> shape;
    
    // Check if it's a scalar
    double scalar_val;
    int int_val;
    if (enif_get_double(env, argv[0], &scalar_val)) {
        data.push_back(static_cast<float>(scalar_val));
        // Scalar has empty shape
    } else if (enif_get_int(env, argv[0], &int_val)) {
        data.push_back(static_cast<float>(int_val));
        // Scalar has empty shape
    } else {
        // Parse nested list
        if (!parse_nested_list(env, argv[0], data, shape)) {
            return enif_make_badarg(env);
        }
    }
    
    // Parse dtype
    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);
        
        // Create array from data
        if (shape.empty()) {
            // Scalar
            array result = array(data[0], dtype);
            return make_array_resource(env, result);
        } else {
            // Create array with specific shape
            array result = array(data.data(), shape, dtype);
            return make_array_resource(env, result);
        }
    } catch (const std::exception& e) {
        return make_error(env, "array_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = sum(a); // Default case
        if (argc == 1) {
            // already initialized with 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = mean(a); // Default case
        if (argc == 1) {
            // already initialized with 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = max(a); // Default case
        if (argc == 1) {
            // already initialized with 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = min(a); // Initialize with default case
        if (argc == 1) {
            // Already initialized
        } 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = var(a); // Initialize with default case
        if (argc == 1) {
            // Already initialized
        } 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = mx::std(a); // Initialize with default case
        if (argc == 1) {
            // Already initialized
        } 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 = mx::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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = transpose(a); // Initialize with default case
        if (argc == 1) {
            // Already initialized
        } 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = squeeze(a); // Initialize with default case
        if (argc == 1) {
            // Already initialized
        } 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
        }
        bool success;
        array arr = get_array_resource_val(env, head, success);
        if (!success) {
            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);
        }
        bool success;
        array arr = get_array_resource_val(env, head, success);
        if (!success) {
            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);
    
    bool success_a, success_b;
    array a = get_array_resource_val(env, argv[0], success_a);
    array b = get_array_resource_val(env, argv[1], success_b);
    if (!success_a || !success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = (argc == 1) ? sort(a) : array({0.0f}); // Initialize with dummy
        if (argc == 1) {
            // Already initialized above
        } 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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = (argc == 1) ? argsort(a) : array({0.0f}); // Initialize with dummy
        if (argc == 1) {
            // Already initialized above
        } 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);
    
    bool success_cond, success_x, success_y;
    array condition = get_array_resource_val(env, argv[0], success_cond);
    array x = get_array_resource_val(env, argv[1], success_x);
    array y = get_array_resource_val(env, argv[2], success_y);
    if (!success_cond || !success_x || !success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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 == complex64) dtype_name = "complex64";
        else dtype_name = "unknown";
        
        return make_ok(env, make_atom(env, dtype_name));
    } catch (const std::exception& e) {
        return make_error(env, "dtype_error");
    }
}

// ==================== RANDOM FUNCTIONS ====================

static ERL_NIF_TERM mlx_random_normal(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 = random::normal(shape, dtype);
        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 != 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 = random::uniform(array(0.0f), array(1.0f), 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 != 3) return enif_make_badarg(env);
    
    double low, high;
    if (!enif_get_double(env, argv[0], &low) ||
        !enif_get_double(env, argv[1], &high)) {
        return enif_make_badarg(env);
    }
    
    auto shape = parse_shape(env, argv[2]);
    if (shape.empty()) return enif_make_badarg(env);
    
    try {
        array result = random::randint(array(static_cast<int>(low)), array(static_cast<int>(high)), shape, int32);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "random_randint_error");
    }
}

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 {
        random::seed(seed);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "random_seed_error");
    }
}

// ==================== ADVANCED MATHEMATICAL FUNCTIONS ====================

static ERL_NIF_TERM mlx_erf(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = erf(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "erf_error");
    }
}

static ERL_NIF_TERM mlx_erfc(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        // erfc(x) = 1 - erf(x)
        array result = subtract(array(1.0f), erf(a));
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "erfc_error");
    }
}

static ERL_NIF_TERM mlx_gamma(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX doesn't have gamma function
        return make_error(env, "gamma_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "gamma_error");
    }
}

static ERL_NIF_TERM mlx_loggamma(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX doesn't have loggamma function
        return make_error(env, "loggamma_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "loggamma_error");
    }
}

static ERL_NIF_TERM mlx_digamma(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX might not have digamma, return error for now
        return make_error(env, "digamma_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "digamma_error");
    }
}

static ERL_NIF_TERM mlx_lgamma(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        // MLX doesn't have lgamma function
        return make_error(env, "lgamma_not_implemented");
    } catch (const std::exception& e) {
        return make_error(env, "lgamma_error");
    }
}

// ==================== DEVICE MANAGEMENT ====================

static ERL_NIF_TERM mlx_set_default_device(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char device_str[32];
    if (!enif_get_atom(env, argv[0], device_str, sizeof(device_str), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        if (strcmp(device_str, "cpu") == 0) {
            set_default_device(Device::cpu);
        } else if (strcmp(device_str, "gpu") == 0) {
            set_default_device(Device::gpu);
        } else {
            return make_error(env, "invalid_device");
        }
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "set_default_device_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);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_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");
    }
}

// Helper function to convert array to nested list
static ERL_NIF_TERM array_to_list_recursive(ErlNifEnv* env, const float* data, const std::vector<int>& shape, int dim, size_t& offset) {
    if (dim == shape.size() - 1) {
        // Base case: create list of values
        ERL_NIF_TERM* terms = new ERL_NIF_TERM[shape[dim]];
        for (int i = 0; i < shape[dim]; i++) {
            terms[i] = enif_make_double(env, data[offset++]);
        }
        ERL_NIF_TERM list = enif_make_list_from_array(env, terms, shape[dim]);
        delete[] terms;
        return list;
    } else {
        // Recursive case: create list of lists
        ERL_NIF_TERM* terms = new ERL_NIF_TERM[shape[dim]];
        for (int i = 0; i < shape[dim]; i++) {
            terms[i] = array_to_list_recursive(env, data, shape, dim + 1, offset);
        }
        ERL_NIF_TERM list = enif_make_list_from_array(env, terms, shape[dim]);
        delete[] terms;
        return list;
    }
}

static ERL_NIF_TERM mlx_to_list(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    bool success_a;
    array a = get_array_resource_val(env, argv[0], success_a);
    if (!success_a) {
        return enif_make_badarg(env);
    }
    
    try {
        // Evaluate the array to ensure data is computed
        eval(a);
        
        // Get shape
        std::vector<int> shape = a.shape();
        
        // Handle scalar case
        if (shape.empty()) {
            return enif_make_double(env, a.item<float>());
        }
        
        // Convert to float32 if needed
        if (a.dtype() != float32) {
            a = astype(a, float32);
            eval(a);
        }
        
        // Get data pointer
        const float* data = a.data<float>();
        
        // Convert to nested list
        size_t offset = 0;
        return array_to_list_recursive(env, data, shape, 0, offset);
        
    } catch (const std::exception& e) {
        return make_error(env, "to_list_error");
    }
}

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

static ERL_NIF_TERM mlx_test_basic(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 0) return enif_make_badarg(env);
    
    try {
        // Simple test: create a 2x2 array of ones and verify it works
        array test = ones({2, 2}, float32);
        eval(test);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "test_basic_failed");
    }
}

static ERL_NIF_TERM mlx_version(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 0) return enif_make_badarg(env);
    
    try {
        // Return MLX Erlang version
        ERL_NIF_TERM version_string = enif_make_string(env, "0.1.0", ERL_NIF_LATIN1);
        return make_ok(env, version_string);
    } catch (const std::exception& e) {
        return make_error(env, "version_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
    {"array", 2, mlx_array, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"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},
    
    // Random functions
    {"random_normal", 2, mlx_random_normal, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"random_uniform", 2, mlx_random_uniform, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"random_randint", 3, mlx_random_randint, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"random_seed", 1, mlx_random_seed, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Advanced mathematical functions
    {"erf", 1, mlx_erf, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"erfc", 1, mlx_erfc, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"gamma", 1, mlx_gamma, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"loggamma", 1, mlx_loggamma, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"digamma", 1, mlx_digamma, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"lgamma", 1, mlx_lgamma, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Device management
    {"set_default_device", 1, mlx_set_default_device, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Evaluation
    {"eval", 1, mlx_eval, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"to_list", 1, mlx_to_list, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Utility functions
    {"test_basic", 0, mlx_test_basic, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"version", 0, mlx_version, 0}
};

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_nif, nif_funcs, load, NULL, upgrade, unload)