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

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

// Resource types for optimizers
static ErlNifResourceType* ARRAY_RESOURCE_TYPE;
static ErlNifResourceType* OPTIMIZER_RESOURCE_TYPE;

struct ArrayResource {
    array arr;
    std::string name;
    ArrayResource(const array& a, const std::string& n = "") : arr(a), name(n) {}
};

struct OptimizerResource {
    std::unique_ptr<opt::Optimizer> optimizer;
    std::string type;
    OptimizerResource(std::unique_ptr<opt::Optimizer> opt, const std::string& t) 
        : optimizer(std::move(opt)), type(t) {}
};

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

// ==================== SGD OPTIMIZER ====================

static ERL_NIF_TERM mlx_sgd_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double learning_rate, momentum, weight_decay;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &momentum) ||
        !enif_get_double(env, argv[2], &weight_decay)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::SGD>(learning_rate, momentum, weight_decay);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "SGD");
        
        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, "sgd_create_error");
    }
}

// ==================== ADAM OPTIMIZER ====================

static ERL_NIF_TERM mlx_adam_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 5) return enif_make_badarg(env);
    
    double learning_rate, beta1, beta2, eps, weight_decay;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &beta1) ||
        !enif_get_double(env, argv[2], &beta2) ||
        !enif_get_double(env, argv[3], &eps) ||
        !enif_get_double(env, argv[4], &weight_decay)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::Adam>(learning_rate, beta1, beta2, eps, weight_decay);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "Adam");
        
        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, "adam_create_error");
    }
}

// ==================== ADAMW OPTIMIZER ====================

static ERL_NIF_TERM mlx_adamw_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 5) return enif_make_badarg(env);
    
    double learning_rate, beta1, beta2, eps, weight_decay;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &beta1) ||
        !enif_get_double(env, argv[2], &beta2) ||
        !enif_get_double(env, argv[3], &eps) ||
        !enif_get_double(env, argv[4], &weight_decay)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::AdamW>(learning_rate, beta1, beta2, eps, weight_decay);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "AdamW");
        
        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, "adamw_create_error");
    }
}

// ==================== RMSPROP OPTIMIZER ====================

static ERL_NIF_TERM mlx_rmsprop_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    double learning_rate, alpha, eps, weight_decay;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &alpha) ||
        !enif_get_double(env, argv[2], &eps) ||
        !enif_get_double(env, argv[3], &weight_decay)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::RMSprop>(learning_rate, alpha, eps, weight_decay);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "RMSprop");
        
        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, "rmsprop_create_error");
    }
}

// ==================== ADAGRAD OPTIMIZER ====================

static ERL_NIF_TERM mlx_adagrad_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double learning_rate, eps, weight_decay;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &eps) ||
        !enif_get_double(env, argv[2], &weight_decay)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::Adagrad>(learning_rate, eps, weight_decay);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "Adagrad");
        
        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, "adagrad_create_error");
    }
}

// ==================== ADADELTA OPTIMIZER ====================

static ERL_NIF_TERM mlx_adadelta_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double learning_rate, rho, eps;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &rho) ||
        !enif_get_double(env, argv[2], &eps)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::Adadelta>(learning_rate, rho, eps);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "Adadelta");
        
        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, "adadelta_create_error");
    }
}

// ==================== LION OPTIMIZER ====================

static ERL_NIF_TERM mlx_lion_create(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    double learning_rate, beta1, beta2, weight_decay;
    if (!enif_get_double(env, argv[0], &learning_rate) ||
        !enif_get_double(env, argv[1], &beta1) ||
        !enif_get_double(env, argv[2], &beta2) ||
        !enif_get_double(env, argv[3], &weight_decay)) {
        return enif_make_badarg(env);
    }
    
    try {
        auto optimizer = std::make_unique<opt::Lion>(learning_rate, beta1, beta2, weight_decay);
        
        OptimizerResource* res = (OptimizerResource*)enif_alloc_resource(
            OPTIMIZER_RESOURCE_TYPE, sizeof(OptimizerResource));
        new(res) OptimizerResource(std::move(optimizer), "Lion");
        
        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, "lion_create_error");
    }
}

// ==================== OPTIMIZER OPERATIONS ====================

static ERL_NIF_TERM mlx_optimizer_update(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    OptimizerResource* opt_res;
    if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
        return enif_make_badarg(env);
    }
    
    array parameters, gradients;
    if (!get_array_resource(env, argv[1], parameters) ||
        !get_array_resource(env, argv[2], gradients)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Apply optimizer update
        auto updates = opt_res->optimizer->apply_gradients({gradients}, {parameters});
        
        if (updates.empty()) {
            return make_error(env, "no_updates");
        }
        
        // Return updated parameters
        return make_array_resource(env, updates[0]);
        
    } catch (const std::exception& e) {
        return make_error(env, "optimizer_update_error");
    }
}

static ERL_NIF_TERM mlx_optimizer_step(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    OptimizerResource* opt_res;
    if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Update optimizer internal state
        opt_res->optimizer->update();
        return make_atom(env, "ok");
        
    } catch (const std::exception& e) {
        return make_error(env, "optimizer_step_error");
    }
}

// ==================== LEARNING RATE SCHEDULERS ====================

static ERL_NIF_TERM mlx_scheduler_step_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    double initial_lr, decay_rate;
    int step, decay_steps;
    if (!enif_get_double(env, argv[0], &initial_lr) ||
        !enif_get_double(env, argv[1], &decay_rate) ||
        !enif_get_int(env, argv[2], &step) ||
        !enif_get_int(env, argv[3], &decay_steps)) {
        return enif_make_badarg(env);
    }
    
    try {
        double current_lr = initial_lr * pow(decay_rate, step / decay_steps);
        return make_ok(env, enif_make_double(env, current_lr));
        
    } catch (const std::exception& e) {
        return make_error(env, "step_decay_error");
    }
}

static ERL_NIF_TERM mlx_scheduler_exponential_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double initial_lr, decay_rate;
    int step;
    if (!enif_get_double(env, argv[0], &initial_lr) ||
        !enif_get_double(env, argv[1], &decay_rate) ||
        !enif_get_int(env, argv[2], &step)) {
        return enif_make_badarg(env);
    }
    
    try {
        double current_lr = initial_lr * pow(decay_rate, step);
        return make_ok(env, enif_make_double(env, current_lr));
        
    } catch (const std::exception& e) {
        return make_error(env, "exponential_decay_error");
    }
}

static ERL_NIF_TERM mlx_scheduler_cosine_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 3) return enif_make_badarg(env);
    
    double initial_lr;
    int step, total_steps;
    if (!enif_get_double(env, argv[0], &initial_lr) ||
        !enif_get_int(env, argv[1], &step) ||
        !enif_get_int(env, argv[2], &total_steps)) {
        return enif_make_badarg(env);
    }
    
    try {
        double progress = static_cast<double>(step) / total_steps;
        double current_lr = initial_lr * 0.5 * (1.0 + cos(M_PI * progress));
        return make_ok(env, enif_make_double(env, current_lr));
        
    } catch (const std::exception& e) {
        return make_error(env, "cosine_decay_error");
    }
}

static ERL_NIF_TERM mlx_scheduler_linear_decay(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 4) return enif_make_badarg(env);
    
    double initial_lr, final_lr;
    int step, total_steps;
    if (!enif_get_double(env, argv[0], &initial_lr) ||
        !enif_get_double(env, argv[1], &final_lr) ||
        !enif_get_int(env, argv[2], &step) ||
        !enif_get_int(env, argv[3], &total_steps)) {
        return enif_make_badarg(env);
    }
    
    try {
        double progress = static_cast<double>(step) / total_steps;
        double current_lr = initial_lr + (final_lr - initial_lr) * progress;
        return make_ok(env, enif_make_double(env, current_lr));
        
    } catch (const std::exception& e) {
        return make_error(env, "linear_decay_error");
    }
}

// ==================== GRADIENT CLIPPING ====================

static ERL_NIF_TERM mlx_clip_grad_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array gradients;
    if (!get_array_resource(env, argv[0], gradients)) {
        return enif_make_badarg(env);
    }
    
    double max_norm;
    if (!enif_get_double(env, argv[1], &max_norm)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Calculate gradient norm
        array grad_norm = sqrt(sum(square(gradients)));
        
        // Clip gradients if norm exceeds max_norm
        array max_norm_arr = array(static_cast<float>(max_norm));
        array clip_coeff = minimum(divide(max_norm_arr, grad_norm), ones_like(grad_norm));
        array clipped_gradients = multiply(gradients, clip_coeff);
        
        return make_array_resource(env, clipped_gradients);
        
    } catch (const std::exception& e) {
        return make_error(env, "clip_grad_norm_error");
    }
}

static ERL_NIF_TERM mlx_clip_grad_value(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    array gradients;
    if (!get_array_resource(env, argv[0], gradients)) {
        return enif_make_badarg(env);
    }
    
    double clip_value;
    if (!enif_get_double(env, argv[1], &clip_value)) {
        return enif_make_badarg(env);
    }
    
    try {
        array clip_val = array(static_cast<float>(clip_value));
        array neg_clip_val = negative(clip_val);
        array clipped_gradients = clip(gradients, neg_clip_val, clip_val);
        
        return make_array_resource(env, clipped_gradients);
        
    } catch (const std::exception& e) {
        return make_error(env, "clip_grad_value_error");
    }
}

// ==================== OPTIMIZER UTILITIES ====================

static ERL_NIF_TERM mlx_get_optimizer_state(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    OptimizerResource* opt_res;
    if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Return optimizer type and learning rate info
        ERL_NIF_TERM type_atom = make_atom(env, opt_res->type.c_str());
        ERL_NIF_TERM state_info = enif_make_tuple2(env,
            make_atom(env, "type"),
            type_atom);
        
        return make_ok(env, state_info);
        
    } catch (const std::exception& e) {
        return make_error(env, "get_optimizer_state_error");
    }
}

static ERL_NIF_TERM mlx_set_learning_rate(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 2) return enif_make_badarg(env);
    
    OptimizerResource* opt_res;
    if (!enif_get_resource(env, argv[0], OPTIMIZER_RESOURCE_TYPE, (void**)&opt_res)) {
        return enif_make_badarg(env);
    }
    
    double new_lr;
    if (!enif_get_double(env, argv[1], &new_lr)) {
        return enif_make_badarg(env);
    }
    
    try {
        // Set new learning rate (simplified - actual implementation would depend on optimizer type)
        opt_res->optimizer->set_learning_rate(new_lr);
        return make_atom(env, "ok");
        
    } catch (const std::exception& e) {
        return make_error(env, "set_learning_rate_error");
    }
}

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

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

// Optimizers NIF function table
static ErlNifFunc nif_funcs[] = {
    // Optimizer creation
    {"sgd_create", 3, mlx_sgd_create, 0},
    {"adam_create", 5, mlx_adam_create, 0},
    {"adamw_create", 5, mlx_adamw_create, 0},
    {"rmsprop_create", 4, mlx_rmsprop_create, 0},
    {"adagrad_create", 3, mlx_adagrad_create, 0},
    {"adadelta_create", 3, mlx_adadelta_create, 0},
    {"lion_create", 4, mlx_lion_create, 0},
    
    // Optimizer operations
    {"optimizer_update", 3, mlx_optimizer_update, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"optimizer_step", 1, mlx_optimizer_step, 0},
    
    // Learning rate schedulers
    {"scheduler_step_decay", 4, mlx_scheduler_step_decay, 0},
    {"scheduler_exponential_decay", 3, mlx_scheduler_exponential_decay, 0},
    {"scheduler_cosine_decay", 3, mlx_scheduler_cosine_decay, 0},
    {"scheduler_linear_decay", 4, mlx_scheduler_linear_decay, 0},
    
    // Gradient clipping
    {"clip_grad_norm", 2, mlx_clip_grad_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"clip_grad_value", 2, mlx_clip_grad_value, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Optimizer utilities
    {"get_optimizer_state", 1, mlx_get_optimizer_state, 0},
    {"set_learning_rate", 2, mlx_set_learning_rate, 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_optimizers_array", array_resource_destructor, flags, tried);
        
    OPTIMIZER_RESOURCE_TYPE = enif_open_resource_type(
        env, NULL, "mlx_optimizer", optimizer_resource_destructor, flags, tried);

    if (!ARRAY_RESOURCE_TYPE || !OPTIMIZER_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_optimizers_nif, nif_funcs, load, NULL, upgrade, unload)