#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/io.h>
#include <memory>
#include <vector>
#include <string>
#include <fstream>
#include <map>

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

// Resource types for I/O operations
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;

struct ArrayResource {
    array arr = array({1.0f}); // Initialize with dummy value
    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);
}

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

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

// ==================== ARRAY SAVING ====================

static ERL_NIF_TERM mlx_save_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    array arr = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[1], arr)) {
        return enif_make_badarg(env);
    }
    
    try {
        io::save(filename, arr);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "save_array_error");
    }
}

static ERL_NIF_TERM mlx_load_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = io::load(filename);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "load_array_error");
    }
}

// ==================== SAFETENSORS FORMAT ====================

static ERL_NIF_TERM mlx_save_safetensors(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    // Parse array dictionary from Erlang
    unsigned int dict_size;
    if (!enif_get_map_size(env, argv[1], &dict_size)) {
        return enif_make_badarg(env);
    }
    
    try {
        std::map<std::string, array> array_dict;
        
        ErlNifMapIterator iter;
        enif_map_iterator_create(env, argv[1], &iter, ERL_NIF_MAP_ITERATOR_FIRST);
        
        ERL_NIF_TERM key, value;
        while (enif_map_iterator_get_pair(env, &iter, &key, &value)) {
            char key_str[256];
            if (!enif_get_string(env, key, key_str, sizeof(key_str), ERL_NIF_LATIN1)) {
                enif_map_iterator_destroy(env, &iter);
                return enif_make_badarg(env);
            }
            
            array arr = array({1.0f}); // Initialize with dummy value
            if (!get_array_resource(env, value, arr)) {
                enif_map_iterator_destroy(env, &iter);
                return enif_make_badarg(env);
            }
            
            array_dict[std::string(key_str)] = arr;
            enif_map_iterator_next(env, &iter);
        }
        enif_map_iterator_destroy(env, &iter);
        
        io::save_safetensors(filename, array_dict);
        return make_atom(env, "ok");
        
    } catch (const std::exception& e) {
        return make_error(env, "save_safetensors_error");
    }
}

static ERL_NIF_TERM mlx_load_safetensors(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto array_dict = io::load_safetensors(filename);
        
        // Convert to Erlang map
        ERL_NIF_TERM result_map = enif_make_new_map(env);
        
        for (const auto& pair : array_dict) {
            ERL_NIF_TERM key = enif_make_string(env, pair.first.c_str(), ERL_NIF_LATIN1);
            auto value_result = make_array_resource(env, pair.second);
            
            // Extract array term from {ok, Array} tuple
            const ERL_NIF_TERM* tuple_elements;
            int tuple_arity;
            ERL_NIF_TERM value_term;
            if (enif_get_tuple(env, value_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
                value_term = tuple_elements[1];
            } else {
                value_term = value_result;
            }
            
            enif_make_map_put(env, result_map, key, value_term, &result_map);
        }
        
        return make_ok(env, result_map);
        
    } catch (const std::exception& e) {
        return make_error(env, "load_safetensors_error");
    }
}

// ==================== NPY FORMAT ====================

static ERL_NIF_TERM mlx_save_npy(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    array arr = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[1], arr)) {
        return enif_make_badarg(env);
    }
    
    try {
        io::save_npy(filename, arr);
        return make_atom(env, "ok");
    } catch (const std::exception& e) {
        return make_error(env, "save_npy_error");
    }
}

static ERL_NIF_TERM mlx_load_npy(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = io::load_npy(filename);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "load_npy_error");
    }
}

// ==================== NPZ FORMAT ====================

static ERL_NIF_TERM mlx_save_npz(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    // Parse array dictionary
    unsigned int dict_size;
    if (!enif_get_map_size(env, argv[1], &dict_size)) {
        return enif_make_badarg(env);
    }
    
    try {
        std::map<std::string, array> array_dict;
        
        ErlNifMapIterator iter;
        enif_map_iterator_create(env, argv[1], &iter, ERL_NIF_MAP_ITERATOR_FIRST);
        
        ERL_NIF_TERM key, value;
        while (enif_map_iterator_get_pair(env, &iter, &key, &value)) {
            char key_str[256];
            if (!enif_get_string(env, key, key_str, sizeof(key_str), ERL_NIF_LATIN1)) {
                enif_map_iterator_destroy(env, &iter);
                return enif_make_badarg(env);
            }
            
            array arr = array({1.0f}); // Initialize with dummy value
            if (!get_array_resource(env, value, arr)) {
                enif_map_iterator_destroy(env, &iter);
                return enif_make_badarg(env);
            }
            
            array_dict[std::string(key_str)] = arr;
            enif_map_iterator_next(env, &iter);
        }
        enif_map_iterator_destroy(env, &iter);
        
        io::save_npz(filename, array_dict);
        return make_atom(env, "ok");
        
    } catch (const std::exception& e) {
        return make_error(env, "save_npz_error");
    }
}

static ERL_NIF_TERM mlx_load_npz(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto array_dict = io::load_npz(filename);
        
        // Convert to Erlang map
        ERL_NIF_TERM result_map = enif_make_new_map(env);
        
        for (const auto& pair : array_dict) {
            ERL_NIF_TERM key = enif_make_string(env, pair.first.c_str(), ERL_NIF_LATIN1);
            auto value_result = make_array_resource(env, pair.second);
            
            // Extract array term from {ok, Array} tuple
            const ERL_NIF_TERM* tuple_elements;
            int tuple_arity;
            ERL_NIF_TERM value_term;
            if (enif_get_tuple(env, value_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
                value_term = tuple_elements[1];
            } else {
                value_term = value_result;
            }
            
            enif_make_map_put(env, result_map, key, value_term, &result_map);
        }
        
        return make_ok(env, result_map);
        
    } catch (const std::exception& e) {
        return make_error(env, "load_npz_error");
    }
}

// ==================== GGUF FORMAT ====================

static ERL_NIF_TERM mlx_save_gguf(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    // Parse array dictionary
    unsigned int dict_size;
    if (!enif_get_map_size(env, argv[1], &dict_size)) {
        return enif_make_badarg(env);
    }
    
    try {
        std::map<std::string, array> array_dict;
        
        ErlNifMapIterator iter;
        enif_map_iterator_create(env, argv[1], &iter, ERL_NIF_MAP_ITERATOR_FIRST);
        
        ERL_NIF_TERM key, value;
        while (enif_map_iterator_get_pair(env, &iter, &key, &value)) {
            char key_str[256];
            if (!enif_get_string(env, key, key_str, sizeof(key_str), ERL_NIF_LATIN1)) {
                enif_map_iterator_destroy(env, &iter);
                return enif_make_badarg(env);
            }
            
            array arr = array({1.0f}); // Initialize with dummy value
            if (!get_array_resource(env, value, arr)) {
                enif_map_iterator_destroy(env, &iter);
                return enif_make_badarg(env);
            }
            
            array_dict[std::string(key_str)] = arr;
            enif_map_iterator_next(env, &iter);
        }
        enif_map_iterator_destroy(env, &iter);
        
        io::save_gguf(filename, array_dict);
        return make_atom(env, "ok");
        
    } catch (const std::exception& e) {
        return make_error(env, "save_gguf_error");
    }
}

static ERL_NIF_TERM mlx_load_gguf(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto array_dict = io::load_gguf(filename);
        
        // Convert to Erlang map
        ERL_NIF_TERM result_map = enif_make_new_map(env);
        
        for (const auto& pair : array_dict) {
            ERL_NIF_TERM key = enif_make_string(env, pair.first.c_str(), ERL_NIF_LATIN1);
            auto value_result = make_array_resource(env, pair.second);
            
            // Extract array term from {ok, Array} tuple
            const ERL_NIF_TERM* tuple_elements;
            int tuple_arity;
            ERL_NIF_TERM value_term;
            if (enif_get_tuple(env, value_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
                value_term = tuple_elements[1];
            } else {
                value_term = value_result;
            }
            
            enif_make_map_put(env, result_map, key, value_term, &result_map);
        }
        
        return make_ok(env, result_map);
        
    } catch (const std::exception& e) {
        return make_error(env, "load_gguf_error");
    }
}

// ==================== ARRAY SERIALIZATION ====================

static ERL_NIF_TERM mlx_array_to_bytes(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array arr = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], arr)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Convert array to bytes (simplified - would use actual MLX serialization)
        eval(arr);  // Ensure array is materialized
        
        auto shape = arr.shape();
        auto dtype = arr.dtype();
        size_t byte_size = arr.nbytes();
        
        // Create byte representation (simplified)
        ErlNifBinary binary;
        if (!enif_alloc_binary(byte_size, &binary)) {
            return make_error(env, "binary_alloc_error");
        }
        
        // Copy array data (simplified - would use proper MLX data access)
        memset(binary.data, 0, byte_size);  // Placeholder
        
        ERL_NIF_TERM binary_term = enif_make_binary(env, &binary);
        return make_ok(env, binary_term);
        
    } catch (const std::exception& e) {
        return make_error(env, "array_to_bytes_error");
    }
}

static ERL_NIF_TERM mlx_bytes_to_array(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    ErlNifBinary binary;
    if (!enif_inspect_binary(env, argv[0], &binary)) {
        return enif_make_badarg(env);
    }
    
    // Parse shape and dtype
    auto shape = parse_shape(env, argv[1]);
    if (shape.empty()) 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 {
        // Create array from bytes (simplified - would use actual MLX deserialization)
        Dtype dtype = parse_dtype(dtype_str);
        array result = zeros(shape, dtype);  // Placeholder
        
        return make_array_resource(env, result);
        
    } catch (const std::exception& e) {
        return make_error(env, "bytes_to_array_error");
    }
}

// ==================== FILE UTILITIES ====================

static ERL_NIF_TERM mlx_file_exists(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    char filename[1024];
    if (!enif_get_string(env, argv[0], filename, sizeof(filename), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    std::ifstream file(filename);
    bool exists = file.good();
    file.close();
    
    return make_ok(env, exists ? make_atom(env, "true") : make_atom(env, "false"));
}

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

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

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

// I/O NIF function table
static ErlNifFunc nif_funcs[] = {
    // Basic array I/O
    {"save_array", 2, mlx_save_array, ERL_NIF_DIRTY_JOB_IO_BOUND},
    {"load_array", 1, mlx_load_array, ERL_NIF_DIRTY_JOB_IO_BOUND},
    
    // SafeTensors format
    {"save_safetensors", 2, mlx_save_safetensors, ERL_NIF_DIRTY_JOB_IO_BOUND},
    {"load_safetensors", 1, mlx_load_safetensors, ERL_NIF_DIRTY_JOB_IO_BOUND},
    
    // NPY format
    {"save_npy", 2, mlx_save_npy, ERL_NIF_DIRTY_JOB_IO_BOUND},
    {"load_npy", 1, mlx_load_npy, ERL_NIF_DIRTY_JOB_IO_BOUND},
    
    // NPZ format
    {"save_npz", 2, mlx_save_npz, ERL_NIF_DIRTY_JOB_IO_BOUND},
    {"load_npz", 1, mlx_load_npz, ERL_NIF_DIRTY_JOB_IO_BOUND},
    
    // GGUF format
    {"save_gguf", 2, mlx_save_gguf, ERL_NIF_DIRTY_JOB_IO_BOUND},
    {"load_gguf", 1, mlx_load_gguf, ERL_NIF_DIRTY_JOB_IO_BOUND},
    
    // Array serialization
    {"array_to_bytes", 1, mlx_array_to_bytes, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"bytes_to_array", 3, mlx_bytes_to_array, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // File utilities
    {"file_exists", 1, mlx_file_exists, ERL_NIF_DIRTY_JOB_IO_BOUND}
};

static int load(ErlNifEnv* env, void** priv_data, ERL_NIF_TERM load_info) {
    ErlNifResourceFlags flags = (ErlNifResourceFlags)(ERL_NIF_RT_CREATE | ERL_NIF_RT_TAKEOVER);
    ErlNifResourceFlags* tried = NULL;
    
    ARRAY_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_io_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_io_nif, nif_funcs, load, NULL, upgrade, unload)