#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/transforms.h>
#include <mlx/compile.h>
#include <memory>
#include <vector>
#include <functional>
#include <map>

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

// Resource types for transforms
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
static ErlNifResourceType* FUNCTION_RESOURCE_TYPE;
static ErlNifResourceType* COMPILED_FUNCTION_RESOURCE_TYPE;

// Function wrapper for storing compiled functions
struct FunctionResource {
    std::function<std::vector<array>(const std::vector<array>&)> fn;
    std::string name;
    FunctionResource(std::function<std::vector<array>(const std::vector<array>&)> f, const std::string& n)
        : fn(f), name(n) {}
};

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

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

// ==================== GRADIENT COMPUTATION ====================

static ERL_NIF_TERM mlx_grad_simple(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) 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);
    }
    
    auto argnums = parse_shape(env, argv[1]);
    if (argnums.empty()) argnums = {0}; // Default to first argument
    
    try {
        std::function<std::vector<array>(const std::vector<array>&)> grad_fn;
        
        if (strcmp(function_name, "square") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.empty()) return {};
                array x = inputs[0];
                array result = square(x);
                return {result};
            }, argnums);
        } else if (strcmp(function_name, "sum_squares") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.empty()) return {};
                array x = inputs[0];
                array result = sum(square(x));
                return {result};
            }, argnums);
        } else if (strcmp(function_name, "mse_loss") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array pred = inputs[0];
                array target = inputs[1];
                array diff = subtract(pred, target);
                array loss = mean(square(diff));
                return {loss};
            }, argnums);
        } else if (strcmp(function_name, "cross_entropy") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array logits = inputs[0];
                array targets = inputs[1];
                array log_probs = log_softmax(logits, -1);
                array loss = negative(mean(sum(multiply(targets, log_probs), -1)));
                return {loss};
            }, argnums);
        } else if (strcmp(function_name, "linear") == 0) {
            grad_fn = grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array x = inputs[0];
                array w = inputs[1];
                array result = matmul(x, w);
                if (inputs.size() > 2) {
                    array b = inputs[2];
                    result = add(result, b);
                }
                return {result};
            }, argnums);
        } else {
            return make_error(env, "unknown_function");
        }
        
        FunctionResource* res = (FunctionResource*)enif_alloc_resource(FUNCTION_RESOURCE_TYPE, sizeof(FunctionResource));
        new(res) FunctionResource(grad_fn, std::string(function_name) + "_grad");
        
        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, "grad_error");
    }
}

// Execute gradient function
static ERL_NIF_TERM mlx_execute_grad(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    FunctionResource* fn_res;
    if (!enif_get_resource(env, argv[0], FUNCTION_RESOURCE_TYPE, (void**)&fn_res)) {
        return enif_make_badarg(env);
    }
    
    // Parse input arrays
    unsigned int array_count;
    if (!enif_get_list_length(env, argv[1], &array_count)) {
        return enif_make_badarg(env);
    }
    
    std::vector<array> inputs;
    inputs.reserve(array_count);
    
    ERL_NIF_TERM head, tail = argv[1];
    for (unsigned int i = 0; i < array_count; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return enif_make_badarg(env);
        }
        array arr;
        if (!get_array_resource(env, head, arr)) {
            return enif_make_badarg(env);
        }
        inputs.push_back(arr);
    }
    
    try {
        std::vector<array> gradients = fn_res->fn(inputs);
        
        // Return gradients as list
        ERL_NIF_TERM grad_list = enif_make_list(env, 0);
        for (int i = gradients.size() - 1; i >= 0; i--) {
            ERL_NIF_TERM grad_term;
            if (make_array_resource(env, gradients[i]) == make_error(env, "array_creation_failed")) {
                return make_error(env, "gradient_creation_failed");
            }
            auto grad_result = make_array_resource(env, gradients[i]);
            // Extract the array term from the {ok, Array} tuple
            const ERL_NIF_TERM* tuple_elements;
            int tuple_arity;
            if (enif_get_tuple(env, grad_result, &tuple_arity, &tuple_elements) && tuple_arity == 2) {
                grad_term = tuple_elements[1];
            } else {
                grad_term = grad_result;
            }
            grad_list = enif_make_list_cell(env, grad_term, grad_list);
        }
        
        return make_ok(env, grad_list);
        
    } catch (const std::exception& e) {
        return make_error(env, "grad_execution_error");
    }
}

// ==================== VALUE AND GRAD ====================

static ERL_NIF_TERM mlx_value_and_grad(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) 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);
    }
    
    auto argnums = parse_shape(env, argv[1]);
    if (argnums.empty()) argnums = {0};
    
    try {
        std::function<std::pair<std::vector<array>, std::vector<array>>(const std::vector<array>&)> value_and_grad_fn;
        
        if (strcmp(function_name, "mse_loss") == 0) {
            value_and_grad_fn = value_and_grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array pred = inputs[0];
                array target = inputs[1];
                array diff = subtract(pred, target);
                array loss = mean(square(diff));
                return {loss};
            }, argnums);
        } else if (strcmp(function_name, "cross_entropy") == 0) {
            value_and_grad_fn = value_and_grad([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array logits = inputs[0];
                array targets = inputs[1];
                array log_probs = log_softmax(logits, -1);
                array loss = negative(mean(sum(multiply(targets, log_probs), -1)));
                return {loss};
            }, argnums);
        } else {
            return make_error(env, "unknown_function");
        }
        
        // For now, return a placeholder indicating success
        return make_ok(env, make_atom(env, "value_and_grad_created"));
        
    } catch (const std::exception& e) {
        return make_error(env, "value_and_grad_error");
    }
}

// ==================== VECTORIZATION (VMAP) ====================

static ERL_NIF_TERM mlx_vmap(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) 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);
    }
    
    auto in_axes = parse_shape(env, argv[1]);
    auto out_axes = parse_shape(env, argv[2]);
    
    try {
        std::function<std::vector<array>(const std::vector<array>&)> vmap_fn;
        
        if (strcmp(function_name, "square") == 0) {
            vmap_fn = vmap([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.empty()) return {};
                array x = inputs[0];
                array result = square(x);
                return {result};
            }, in_axes, out_axes);
        } else if (strcmp(function_name, "matmul") == 0) {
            vmap_fn = vmap([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array a = inputs[0];
                array b = inputs[1];
                array result = matmul(a, b);
                return {result};
            }, in_axes, out_axes);
        } else {
            return make_error(env, "unknown_function");
        }
        
        FunctionResource* res = (FunctionResource*)enif_alloc_resource(FUNCTION_RESOURCE_TYPE, sizeof(FunctionResource));
        new(res) FunctionResource(vmap_fn, std::string(function_name) + "_vmap");
        
        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, "vmap_error");
    }
}

// ==================== COMPILATION ====================

static ERL_NIF_TERM mlx_compile(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_linear") == 0) {
            compiled_fn = compile([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 3) return {};
                array x = inputs[0];
                array w = inputs[1];
                array b = inputs[2];
                array result = add(matmul(x, w), b);
                return {result};
            });
        } else 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];
                array k = inputs[1];
                array v = inputs[2];
                
                array scores = matmul(q, transpose(k, {-2, -1}));
                array attn = softmax(scores, -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];
                array w1 = inputs[1];
                array 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, "compile_error");
    }
}

// Execute compiled function
static ERL_NIF_TERM mlx_execute_compiled(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    CompiledFunctionResource* fn_res;
    if (!enif_get_resource(env, argv[0], COMPILED_FUNCTION_RESOURCE_TYPE, (void**)&fn_res)) {
        return enif_make_badarg(env);
    }
    
    // Parse input arrays
    unsigned int array_count;
    if (!enif_get_list_length(env, argv[1], &array_count)) {
        return enif_make_badarg(env);
    }
    
    std::vector<array> inputs;
    inputs.reserve(array_count);
    
    ERL_NIF_TERM head, tail = argv[1];
    for (unsigned int i = 0; i < array_count; i++) {
        if (!enif_get_list_cell(env, tail, &head, &tail)) {
            return enif_make_badarg(env);
        }
        array arr;
        if (!get_array_resource(env, head, arr)) {
            return enif_make_badarg(env);
        }
        inputs.push_back(arr);
    }
    
    try {
        std::vector<array> results = fn_res->compiled_fn(inputs);
        
        if (results.empty()) {
            return make_error(env, "no_results");
        }
        
        // Return first result for simplicity
        return make_array_resource(env, results[0]);
        
    } catch (const std::exception& e) {
        return make_error(env, "compiled_execution_error");
    }
}

// ==================== CHECKPOINT ====================

static ERL_NIF_TERM mlx_checkpoint(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>&)> checkpoint_fn;
        
        if (strcmp(function_name, "memory_efficient_attention") == 0) {
            checkpoint_fn = checkpoint([](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 3) return {};
                array q = inputs[0];
                array k = inputs[1];
                array v = inputs[2];
                
                array scores = matmul(q, transpose(k, {-2, -1}));
                array attn = softmax(scores, -1);
                array result = matmul(attn, v);
                return {result};
            });
        } else {
            return make_error(env, "unknown_function");
        }
        
        FunctionResource* res = (FunctionResource*)enif_alloc_resource(FUNCTION_RESOURCE_TYPE, sizeof(FunctionResource));
        new(res) FunctionResource(checkpoint_fn, std::string(function_name) + "_checkpoint");
        
        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, "checkpoint_error");
    }
}

// ==================== CUSTOM TRANSFORMS ====================

static ERL_NIF_TERM mlx_custom_transform(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    char transform_type[64], function_name[64];
    if (!enif_get_atom(env, argv[0], transform_type, sizeof(transform_type), ERL_NIF_LATIN1) ||
        !enif_get_atom(env, argv[1], function_name, sizeof(function_name), ERL_NIF_LATIN1)) {
        return enif_make_badarg(env);
    }
    
    try {
        std::function<std::vector<array>(const std::vector<array>&)> transform_fn;
        
        if (strcmp(transform_type, "grad_and_compile") == 0) {
            // Compose grad and compile transforms
            auto base_fn = [](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.size() < 2) return {};
                array pred = inputs[0];
                array target = inputs[1];
                array loss = mean(square(subtract(pred, target)));
                return {loss};
            };
            
            auto grad_fn = grad(base_fn, {0});
            transform_fn = compile(grad_fn);
            
        } else if (strcmp(transform_type, "vmap_and_grad") == 0) {
            // Compose vmap and grad transforms
            auto base_fn = [](const std::vector<array>& inputs) -> std::vector<array> {
                if (inputs.empty()) return {};
                array x = inputs[0];
                array result = sum(square(x));
                return {result};
            };
            
            auto vmap_fn = vmap(base_fn, {0}, {0});
            transform_fn = grad(vmap_fn, {0});
            
        } else {
            return make_error(env, "unknown_transform");
        }
        
        FunctionResource* res = (FunctionResource*)enif_alloc_resource(FUNCTION_RESOURCE_TYPE, sizeof(FunctionResource));
        new(res) FunctionResource(transform_fn, std::string(transform_type) + "_" + 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, "custom_transform_error");
    }
}

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

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

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

// Transform NIF function table
static ErlNifFunc nif_funcs[] = {
    // Gradient computation
    {"grad_simple", 2, mlx_grad_simple, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"execute_grad", 2, mlx_execute_grad, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"value_and_grad", 2, mlx_value_and_grad, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Vectorization
    {"vmap", 3, mlx_vmap, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Compilation
    {"compile", 1, mlx_compile, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"execute_compiled", 2, mlx_execute_compiled, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Memory optimization
    {"checkpoint", 1, mlx_checkpoint, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Custom transforms
    {"custom_transform", 2, mlx_custom_transform, ERL_NIF_DIRTY_JOB_CPU_BOUND}
};

static int load(ErlNifEnv* env, void** priv_data, ERL_NIF_TERM load_info) {
    ErlNifResourceFlags flags = (ErlNifResourceFlags)(ERL_NIF_RT_CREATE | ERL_NIF_RT_TAKEOVER);
    ErlNifResourceFlags* tried = NULL;
    
    ARRAY_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_transforms_array", array_resource_destructor, flags, tried);
        
    FUNCTION_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_function", function_resource_destructor, flags, tried);
        
    COMPILED_FUNCTION_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_compiled_function", compiled_function_resource_destructor, flags, tried);

    if (!ARRAY_RESOURCE_TYPE || !FUNCTION_RESOURCE_TYPE || !COMPILED_FUNCTION_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_transforms_nif, nif_funcs, load, NULL, upgrade, unload)