#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/fast.h>
#include <mlx/compile.h>
#include <mlx/random.h>
#include <mlx/linalg.h>
// #include <mlx/nn.h>  // Not available in C++ MLX
#include <memory>
#include <vector>
#include <map>
#include <string>
#include <iostream>
#include <functional>
#include <future>
#include <thread>
#include <atomic>
#include <mutex>

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

// Advanced resource types for sophisticated MLX operations
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
static ErlNifResourceType* STREAM_RESOURCE_TYPE;
static ErlNifResourceType* COMPILED_FUNCTION_RESOURCE_TYPE;
static ErlNifResourceType* GRADIENT_FUNCTION_RESOURCE_TYPE;
static ErlNifResourceType* ATTENTION_CACHE_RESOURCE_TYPE;

// Thread pool for async operations
static std::unique_ptr<std::vector<std::thread>> worker_threads;
static std::atomic<bool> shutdown_flag{false};

// Resource wrappers with advanced capabilities
struct ArrayResource {
    array arr;
    std::string name;
    bool is_parameter = false;
    std::shared_ptr<std::vector<array>> gradient_tape;
    
    ArrayResource(const array& a, const std::string& n = "") 
        : arr(a), name(n), gradient_tape(std::make_shared<std::vector<array>>()) {}
};

struct StreamResource {
    Stream stream;
    StreamResource(const Stream& s) : stream(s) {}
};

struct CompiledFunctionResource {
    std::function<std::vector<array>(const std::vector<array>&)> compiled_fn;
    std::string signature;
    CompiledFunctionResource(std::function<std::vector<array>(const std::vector<array>&)> fn, const std::string& sig)
        : compiled_fn(fn), signature(sig) {}
};

struct GradientFunctionResource {
    std::function<std::vector<array>(const std::vector<array>&)> grad_fn;
    std::vector<int> argnums;
    GradientFunctionResource(std::function<std::vector<array>(const std::vector<array>&)> fn, const std::vector<int>& args)
        : grad_fn(fn), argnums(args) {}
};

struct AttentionCacheResource {
    std::map<std::string, array> kv_cache;
    int max_seq_length;
    int num_heads;
    int head_dim;
    
    AttentionCacheResource(int max_seq, int heads, int dim) 
        : max_seq_length(max_seq), num_heads(heads), head_dim(dim) {}
};

// Performance monitoring
struct PerformanceMetrics {
    std::atomic<size_t> operations_count{0};
    std::atomic<size_t> memory_allocated{0};
    std::atomic<size_t> gpu_operations{0};
    std::atomic<size_t> cpu_operations{0};
    std::mutex metrics_mutex;
    std::map<std::string, double> operation_times;
};

static PerformanceMetrics perf_metrics;

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

// Advanced dtype parsing with all MLX types
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; // default
}

// 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;
}

// Advanced array creation with unified memory optimization
static ERL_NIF_TERM mlx_create_advanced(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);

    auto shape = parse_shape(env, argv[0]);
    if (shape.empty()) return enif_make_badarg(env);

    char dtype_str[32], init_type[32], name[128];
    if (!enif_get_atom(env, argv[1], dtype_str, sizeof(dtype_str), ERL_NIF_LATIN1) ||
        !enif_get_atom(env, argv[2], init_type, sizeof(init_type), ERL_NIF_LATIN1) ||
        !enif_get_string(env, argv[3], name, sizeof(name), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }

    Dtype dtype = parse_dtype(dtype_str);

    try {
        array result;
        
        // Advanced initialization patterns
        if (strcmp(init_type, "zeros") == 0) {
            result = zeros(shape, dtype);
        } else if (strcmp(init_type, "ones") == 0) {
            result = ones(shape, dtype);
        } else if (strcmp(init_type, "eye") == 0) {
            result = eye(shape[0], dtype);
        } else if (strcmp(init_type, "xavier_uniform") == 0) {
            double fan_in = shape.size() > 1 ? shape[1] : 1;
            double fan_out = shape[0];
            double gain = sqrt(6.0 / (fan_in + fan_out));
            result = random::uniform(-gain, gain, shape, dtype);
        } else if (strcmp(init_type, "xavier_normal") == 0) {
            double fan_in = shape.size() > 1 ? shape[1] : 1;
            double fan_out = shape[0];
            double std_dev = sqrt(2.0 / (fan_in + fan_out));
            result = random::normal(0.0, std_dev, shape, dtype);
        } else if (strcmp(init_type, "kaiming_uniform") == 0) {
            double fan_in = shape.size() > 1 ? shape[1] : 1;
            double gain = sqrt(6.0 / fan_in);
            result = random::uniform(-gain, gain, shape, dtype);
        } else if (strcmp(init_type, "kaiming_normal") == 0) {
            double fan_in = shape.size() > 1 ? shape[1] : 1;
            double std_dev = sqrt(2.0 / fan_in);
            result = random::normal(0.0, std_dev, shape, dtype);
        } else {
            return make_error(env, "invalid_init_type");
        }

        // Optimize for unified memory
        result = ensure_row_contiguous(result);
        
        ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
        new(res) ArrayResource(result, std::string(name));
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        perf_metrics.operations_count++;
        
        return make_ok(env, term);
    } catch (const std::exception& e) {
        return make_error(env, "mlx_advanced_create_error");
    }
}

// High-performance linear algebra operations
static ERL_NIF_TERM mlx_advanced_linalg(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);

    ArrayResource* a_res;
    if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&a_res)) {
        return enif_make_badarg(env);
    }

    char operation[32];
    if (!enif_get_atom(env, argv[1], operation, sizeof(operation), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }

    try {
        array result;
        auto start_time = std::chrono::high_resolution_clock::now();

        if (strcmp(operation, "svd") == 0) {
            auto svd_result = linalg::svd(a_res->arr);
            // Return U, S, Vt as tuple
            ArrayResource* u_res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
            ArrayResource* s_res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
            ArrayResource* vt_res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
            
            new(u_res) ArrayResource(std::get<0>(svd_result));
            new(s_res) ArrayResource(std::get<1>(svd_result));
            new(vt_res) ArrayResource(std::get<2>(svd_result));
            
            ERL_NIF_TERM u_term = enif_make_resource(env, u_res);
            ERL_NIF_TERM s_term = enif_make_resource(env, s_res);
            ERL_NIF_TERM vt_term = enif_make_resource(env, vt_res);
            
            enif_release_resource(u_res);
            enif_release_resource(s_res);
            enif_release_resource(vt_res);
            
            return make_ok(env, enif_make_tuple3(env, u_term, s_term, vt_term));
            
        } else if (strcmp(operation, "qr") == 0) {
            auto qr_result = linalg::qr(a_res->arr);
            result = std::get<0>(qr_result); // Return Q for now
            
        } else if (strcmp(operation, "cholesky") == 0) {
            result = linalg::cholesky(a_res->arr);
            
        } else if (strcmp(operation, "inv") == 0) {
            result = linalg::inv(a_res->arr);
            
        } else if (strcmp(operation, "norm") == 0) {
            result = linalg::norm(a_res->arr);
            
        } else if (strcmp(operation, "det") == 0) {
            result = linalg::det(a_res->arr);
            
        } else {
            return make_error(env, "invalid_linalg_operation");
        }

        auto end_time = std::chrono::high_resolution_clock::now();
        auto duration = std::chrono::duration_cast<std::chrono::microseconds>(end_time - start_time);
        
        {
            std::lock_guard<std::mutex> lock(perf_metrics.metrics_mutex);
            perf_metrics.operation_times[operation] = duration.count() / 1000.0; // ms
        }

        ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
        new(res) ArrayResource(result);
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        return make_ok(env, term);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_linalg_error");
    }
}

// Advanced neural network operations with Apple Silicon optimization
static ERL_NIF_TERM mlx_transformer_attention(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 6) return enif_make_badarg(env);

    ArrayResource *q_res, *k_res, *v_res;
    if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&q_res) ||
        !enif_get_resource(env, argv[1], ARRAY_RESOURCE_TYPE, (void**)&k_res) ||
        !enif_get_resource(env, argv[2], ARRAY_RESOURCE_TYPE, (void**)&v_res)) {
        return enif_make_badarg(env);
    }

    double scale;
    int is_causal;
    if (!enif_get_double(env, argv[3], &scale) ||
        !enif_get_int(env, argv[4], &is_causal)) {
        return enif_make_badarg(env);
    }

    try {
        array q = q_res->arr;
        array k = k_res->arr;
        array v = v_res->arr;

        // Scaled dot-product attention optimized for Apple Silicon
        array scores = matmul(q, transpose(k, {-2, -1})) * array(scale);
        
        if (is_causal) {
            // Create causal mask
            auto seq_len = q.shape(-2);
            array mask = triu(ones({seq_len, seq_len}, float32), 1) * (-1e9);
            scores = scores + mask;
        }

        array attn_weights = softmax(scores, -1);
        array result = matmul(attn_weights, v);

        // Use Metal GPU for computation
        result = result.astype(q.dtype());
        eval(result);

        ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
        new(res) ArrayResource(result);
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        perf_metrics.gpu_operations++;
        
        return make_ok(env, term);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_attention_error");
    }
}

// RoPE (Rotary Position Embedding) implementation
static ERL_NIF_TERM mlx_rope_embedding(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);

    ArrayResource* x_res;
    if (!enif_get_resource(env, argv[0], ARRAY_RESOURCE_TYPE, (void**)&x_res)) {
        return enif_make_badarg(env);
    }

    int dims, offset;
    if (!enif_get_int(env, argv[1], &dims) ||
        !enif_get_int(env, argv[2], &offset)) {
        return enif_make_badarg(env);
    }

    try {
        array x = x_res->arr;
        
        // RoPE implementation optimized for Apple Silicon
        auto shape = x.shape();
        int seq_len = shape[-2];
        int hidden_dim = shape[-1];
        
        // Create position encodings
        array positions = arange(offset, offset + seq_len, float32);
        positions = reshape(positions, {seq_len, 1});
        
        // Create frequency bands
        array freqs = exp(arange(0, dims, 2, float32) * (-log(10000.0) / dims));
        freqs = reshape(freqs, {1, dims / 2});
        
        array angles = matmul(positions, freqs);
        array cos_vals = cos(angles);
        array sin_vals = sin(angles);
        
        // Apply RoPE rotation
        auto x_even = x.index({Slice(), Slice(), Slice(None, None, 2)});
        auto x_odd = x.index({Slice(), Slice(), Slice(1, None, 2)});
        
        array rotated_even = x_even * cos_vals - x_odd * sin_vals;
        array rotated_odd = x_even * sin_vals + x_odd * cos_vals;
        
        // Interleave results
        std::vector<array> interleaved;
        for (int i = 0; i < dims / 2; i++) {
            interleaved.push_back(rotated_even.index({Slice(), Slice(), i}));
            interleaved.push_back(rotated_odd.index({Slice(), Slice(), i}));
        }
        
        array result = stack(interleaved, -1);
        eval(result);

        ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
        new(res) ArrayResource(result);
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        return make_ok(env, term);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_rope_error");
    }
}

// Function compilation for performance optimization
static ERL_NIF_TERM mlx_compile_function(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);

    char function_name[64];
    if (!enif_get_atom(env, argv[0], function_name, sizeof(function_name), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }

    try {
        std::function<std::vector<array>(const std::vector<array>&)> compiled_fn;

        if (strcmp(function_name, "fused_attention") == 0) {
            compiled_fn = compile([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() != 3) return {};
                
                array q = inputs[0], k = inputs[1], v = inputs[2];
                array scores = matmul(q, transpose(k, {-2, -1}));
                array attn = softmax(scores / sqrt(array(static_cast<float>(k.shape(-1)))), -1);
                array result = matmul(attn, v);
                
                return {result};
            });
        } else if (strcmp(function_name, "fused_mlp") == 0) {
            compiled_fn = compile([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() != 3) return {};
                
                array x = inputs[0], w1 = inputs[1], w2 = inputs[2];
                array hidden = gelu(matmul(x, w1));
                array result = matmul(hidden, w2);
                
                return {result};
            });
        } else {
            return make_error(env, "unknown_function");
        }

        CompiledFunctionResource* res = (CompiledFunctionResource*)enif_alloc_resource(
            COMPILED_FUNCTION_RESOURCE_TYPE, sizeof(CompiledFunctionResource));
        new(res) CompiledFunctionResource(compiled_fn, std::string(function_name));
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        return make_ok(env, term);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_compile_error");
    }
}

// Automatic differentiation with advanced gradient computation
static ERL_NIF_TERM mlx_create_gradient_function(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);

    char function_type[64];
    if (!enif_get_atom(env, argv[0], function_type, sizeof(function_type), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }

    // Parse argnums list
    auto argnums_list = parse_shape(env, argv[1]);

    try {
        std::function<std::vector<array>(const std::vector<array>&)> grad_fn;

        if (strcmp(function_type, "mse_loss") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() != 2) return {};
                array pred = inputs[0], target = inputs[1];
                array loss = mean(square(pred - target));
                return {loss};
            }, argnums_list);
        } else if (strcmp(function_type, "cross_entropy") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() != 2) return {};
                array logits = inputs[0], targets = inputs[1];
                array log_probs = log_softmax(logits, -1);
                array loss = -mean(sum(targets * log_probs, -1));
                return {loss};
            }, argnums_list);
        } else {
            return make_error(env, "unknown_function_type");
        }

        GradientFunctionResource* res = (GradientFunctionResource*)enif_alloc_resource(
            GRADIENT_FUNCTION_RESOURCE_TYPE, sizeof(GradientFunctionResource));
        new(res) GradientFunctionResource(grad_fn, argnums_list);
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        return make_ok(env, term);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_grad_error");
    }
}

// Advanced memory management with unified memory optimization
static ERL_NIF_TERM mlx_memory_stats(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 0) return enif_make_badarg(env);

    try {
        ERL_NIF_TERM stats = enif_make_tuple4(env,
            enif_make_tuple2(env, make_atom(env, "operations_count"), 
                           enif_make_ulong(env, perf_metrics.operations_count.load())),
            enif_make_tuple2(env, make_atom(env, "memory_allocated"), 
                           enif_make_ulong(env, perf_metrics.memory_allocated.load())),
            enif_make_tuple2(env, make_atom(env, "gpu_operations"), 
                           enif_make_ulong(env, perf_metrics.gpu_operations.load())),
            enif_make_tuple2(env, make_atom(env, "cpu_operations"), 
                           enif_make_ulong(env, perf_metrics.cpu_operations.load()))
        );
        
        return make_ok(env, stats);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_memory_stats_error");
    }
}

// Advanced streaming operations for real-time processing
static ERL_NIF_TERM mlx_create_stream(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 {
        Device device = Device::gpu; // Default to GPU for Apple Silicon
        if (strcmp(device_str, "cpu") == 0) device = Device::cpu;
        else if (strcmp(device_str, "gpu") == 0) device = Device::gpu;

        Stream stream = new_stream(device);
        
        StreamResource* res = (StreamResource*)enif_alloc_resource(STREAM_RESOURCE_TYPE, sizeof(StreamResource));
        new(res) StreamResource(stream);
        
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        
        return make_ok(env, term);
        
    } catch (const std::exception& e) {
        return make_error(env, "mlx_stream_error");
    }
}

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

static void stream_resource_destructor(ErlNifEnv* env, void* obj) {
    StreamResource* res = (StreamResource*)obj;
    res->~StreamResource();
}

static void compiled_function_resource_destructor(ErlNifEnv* env, void* obj) {
    CompiledFunctionResource* res = (CompiledFunctionResource*)obj;
    res->~CompiledFunctionResource();
}

static void gradient_function_resource_destructor(ErlNifEnv* env, void* obj) {
    GradientFunctionResource* res = (GradientFunctionResource*)obj;
    res->~GradientFunctionResource();
}

static void attention_cache_resource_destructor(ErlNifEnv* env, void* obj) {
    AttentionCacheResource* res = (AttentionCacheResource*)obj;
    res->~AttentionCacheResource();
}

// Advanced NIF function table with sophisticated operations
static ErlNifFunc nif_funcs[] = {
    // Advanced array operations
    {"create_advanced", 4, mlx_create_advanced, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // High-performance linear algebra
    {"advanced_linalg", 3, mlx_advanced_linalg, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Neural network operations
    {"transformer_attention", 6, mlx_transformer_attention, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"rope_embedding", 3, mlx_rope_embedding, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Function compilation and optimization
    {"compile_function", 1, mlx_compile_function, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"create_gradient_function", 2, mlx_create_gradient_function, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Memory and performance
    {"memory_stats", 0, mlx_memory_stats, 0},
    
    // Streaming operations
    {"create_stream", 1, mlx_create_stream, 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_advanced_array", array_resource_destructor, flags, tried);
    
    STREAM_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_advanced_stream", stream_resource_destructor, flags, tried);
        
    COMPILED_FUNCTION_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_compiled_function", compiled_function_resource_destructor, flags, tried);
        
    GRADIENT_FUNCTION_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_gradient_function", gradient_function_resource_destructor, flags, tried);
        
    ATTENTION_CACHE_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_attention_cache", attention_cache_resource_destructor, flags, tried);

    if (!ARRAY_RESOURCE_TYPE || !STREAM_RESOURCE_TYPE || 
        !COMPILED_FUNCTION_RESOURCE_TYPE || !GRADIENT_FUNCTION_RESOURCE_TYPE ||
        !ATTENTION_CACHE_RESOURCE_TYPE) {
        return 1;
    }

    // Initialize thread pool for async operations
    worker_threads = std::make_unique<std::vector<std::thread>>();
    
    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) {
    shutdown_flag = true;
    if (worker_threads) {
        for (auto& thread : *worker_threads) {
            if (thread.joinable()) {
                thread.join();
            }
        }
        worker_threads.reset();
    }
}

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