#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 managing MLX objects
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
static ErlNifResourceType* STREAM_RESOURCE_TYPE;

// Resource wrappers
struct ArrayResource {
    array arr;
    ArrayResource(const array& a) : arr(a) {}
};

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

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

// Array creation functions
static ERL_NIF_TERM mlx_zeros(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) {
        return enif_make_badarg(env);
    }

    // Parse shape
    unsigned int shape_len;
    if (!enif_get_list_length(env, argv[0], &shape_len)) {
        return enif_make_badarg(env);
    }

    std::vector<int> shape_vec(shape_len);
    ERL_NIF_TERM head, tail = argv[0];
    
    for (unsigned int i = 0; i < shape_len; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return enif_make_badarg(env);
        }
        if (!enif_get_int(env, head, &shape_vec[i])) {
            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);
    }

    Dtype dtype = float32;  // Default initialization
    if (strcmp(dtype_str, "float32") == 0) dtype = float32;
    else if (strcmp(dtype_str, "float16") == 0) dtype = float16;
    else if (strcmp(dtype_str, "bfloat16") == 0) dtype = bfloat16;
    else if (strcmp(dtype_str, "int32") == 0) dtype = int32;
    else if (strcmp(dtype_str, "int16") == 0) dtype = int16;
    else if (strcmp(dtype_str, "int8") == 0) dtype = int8;
    else if (strcmp(dtype_str, "uint32") == 0) dtype = uint32;
    else if (strcmp(dtype_str, "uint16") == 0) dtype = uint16;
    else if (strcmp(dtype_str, "uint8") == 0) dtype = uint8;
    else if (strcmp(dtype_str, "bool") == 0) dtype = bool_;
    else return make_error(env, "invalid_dtype");

    try {
        array result = zeros(shape_vec, dtype);
        
        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_error");
    }
}

static ERL_NIF_TERM mlx_ones(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) {
        return enif_make_badarg(env);
    }

    // Parse shape
    unsigned int shape_len;
    if (!enif_get_list_length(env, argv[0], &shape_len)) {
        return enif_make_badarg(env);
    }

    std::vector<int> shape_vec(shape_len);
    ERL_NIF_TERM head, tail = argv[0];
    
    for (unsigned int i = 0; i < shape_len; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return enif_make_badarg(env);
        }
        if (!enif_get_int(env, head, &shape_vec[i])) {
            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);
    }

    Dtype dtype = float32;  // Default initialization
    if (strcmp(dtype_str, "float32") == 0) dtype = float32;
    else if (strcmp(dtype_str, "float16") == 0) dtype = float16;
    else if (strcmp(dtype_str, "bfloat16") == 0) dtype = bfloat16;
    else if (strcmp(dtype_str, "int32") == 0) dtype = int32;
    else if (strcmp(dtype_str, "int16") == 0) dtype = int16;
    else if (strcmp(dtype_str, "int8") == 0) dtype = int8;
    else if (strcmp(dtype_str, "uint32") == 0) dtype = uint32;
    else if (strcmp(dtype_str, "uint16") == 0) dtype = uint16;
    else if (strcmp(dtype_str, "uint8") == 0) dtype = uint8;
    else if (strcmp(dtype_str, "bool") == 0) dtype = bool_;
    else return make_error(env, "invalid_dtype");

    try {
        array result = ones(shape_vec, dtype);
        
        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_error");
    }
}

// Array operations
static ERL_NIF_TERM mlx_add(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) {
        return enif_make_badarg(env);
    }

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

    try {
        array result = add(a_res->arr, b_res->arr);
        
        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_error");
    }
}

static ERL_NIF_TERM mlx_multiply(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) {
        return enif_make_badarg(env);
    }

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

    try {
        array result = multiply(a_res->arr, b_res->arr);
        
        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_error");
    }
}

static ERL_NIF_TERM mlx_matmul(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) {
        return enif_make_badarg(env);
    }

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

    try {
        array result = matmul(a_res->arr, b_res->arr);
        
        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_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);
    }

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

    try {
        const std::vector<int>& shape = res->arr.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, "mlx_error");
    }
}

static ERL_NIF_TERM mlx_eval(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) {
        return enif_make_badarg(env);
    }

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

    try {
        eval(res->arr);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "mlx_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) {
        std::cerr << "MLX device error: " << e.what() << std::endl;
        return make_error(env, "mlx_error");
    }
}

// Test function to verify MLX is working
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 {
        // Test basic MLX functionality
        array a = ones({2, 3}, float32);
        array b = zeros({2, 3}, float32);
        array c = add(a, b);
        
        eval(c);  // Force evaluation
        
        auto shape = c.shape();
        auto dtype = c.dtype();
        
        // Return shape and dtype info
        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);
        }
        
        ERL_NIF_TERM dtype_atom;
        if (dtype == float32) dtype_atom = make_atom(env, "float32");
        else if (dtype == int32) dtype_atom = make_atom(env, "int32");
        else dtype_atom = make_atom(env, "unknown");
        
        ERL_NIF_TERM result = enif_make_tuple3(env, 
            make_atom(env, "test_passed"),
            shape_list,
            dtype_atom);
        
        return make_ok(env, result);
        
    } catch (const std::exception& e) {
        std::cerr << "MLX test error: " << e.what() << std::endl;
        return make_error(env, "mlx_test_failed");
    }
}

// Get MLX version info
static ERL_NIF_TERM mlx_version(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 0) {
        return enif_make_badarg(env);
    }
    
    try {
        // Create a simple test to verify MLX is loaded
        array test = ones({1}, float32);
        eval(test);
        
        return make_ok(env, make_atom(env, "mlx_loaded"));
    } catch (const std::exception& e) {
        std::cerr << "MLX version check error: " << e.what() << std::endl;
        return make_error(env, "mlx_not_loaded");
    }
}

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

// NIF function table using dirty schedulers for CPU-intensive operations
static ErlNifFunc nif_funcs[] = {
    {"zeros", 2, mlx_zeros, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"ones", 2, mlx_ones, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"add", 2, mlx_add, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"multiply", 2, mlx_multiply, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"matmul", 2, mlx_matmul, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"shape", 1, mlx_shape, 0},
    {"eval", 1, mlx_eval, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"set_default_device", 1, mlx_set_default_device, 0},
    {"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) {
    // Initialize resource types
    ErlNifResourceFlags flags = (ErlNifResourceFlags)(ERL_NIF_RT_CREATE | ERL_NIF_RT_TAKEOVER);
    ErlNifResourceFlags* tried = NULL;
    
    ARRAY_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_array", array_resource_destructor,
        flags, tried);
    
    STREAM_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_stream", stream_resource_destructor,
        flags, tried);

    if (ARRAY_RESOURCE_TYPE == NULL || STREAM_RESOURCE_TYPE == NULL) {
        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)