#include <erl_nif.h>
#include <mlx/mlx.h>
#include <mlx/ops.h>
#include <mlx/array.h>
#include <mlx/linalg.h>
#include <memory>
#include <vector>

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

// Resource types for linear algebra 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);
}

// Create tuple of array resources
static ERL_NIF_TERM make_array_tuple(ErlNifEnv* env, const std::vector<array>& arrays) {
    std::vector<ERL_NIF_TERM> terms;
    for (const auto& arr : arrays) {
        ArrayResource* res = (ArrayResource*)enif_alloc_resource(ARRAY_RESOURCE_TYPE, sizeof(ArrayResource));
        new(res) ArrayResource(arr);
        ERL_NIF_TERM term = enif_make_resource(env, res);
        enif_release_resource(res);
        terms.push_back(term);
    }
    
    ERL_NIF_TERM tuple;
    if (arrays.size() == 2) {
        tuple = enif_make_tuple2(env, terms[0], terms[1]);
    } else if (arrays.size() == 3) {
        tuple = enif_make_tuple3(env, terms[0], terms[1], terms[2]);
    } else {
        // For arbitrary number of arrays, create a list
        ERL_NIF_TERM list = enif_make_list_from_array(env, terms.data(), terms.size());
        tuple = list;
    }
    
    return make_ok(env, tuple);
}

// ==================== SINGULAR VALUE DECOMPOSITION ====================

static ERL_NIF_TERM mlx_svd(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    bool full_matrices = true;
    if (argc > 1) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        full_matrices = (strcmp(atom_name, "true") == 0);
    }
    
    try {
        auto result = linalg::svd(a, full_matrices, {});
        // result is already a std::vector<array>
        return make_array_tuple(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "svd_error");
    }
}

// ==================== EIGENVALUE DECOMPOSITION ====================

static ERL_NIF_TERM mlx_eig(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        // TODO: MLX only has eigh (hermitian), not general eig
        return make_error(env, "eig_not_available_use_eigh");
    } catch (const std::exception& e) {
        return make_error(env, "eig_error");
    }
}

static ERL_NIF_TERM mlx_eigvals(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        // TODO: MLX only has eigvalsh (hermitian), not general eigvals
        return make_error(env, "eigvals_not_available_use_eigvalsh");
    } catch (const std::exception& e) {
        return make_error(env, "eigvals_error");
    }
}

static ERL_NIF_TERM mlx_eigh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    std::string uplo = "L";
    if (argc > 1) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        uplo = std::string(atom_name);
    }
    
    try {
        auto result = linalg::eigh(a, uplo);
        std::vector<array> arrays = {std::get<0>(result), std::get<1>(result)};
        return make_array_tuple(env, arrays);
    } catch (const std::exception& e) {
        return make_error(env, "eigh_error");
    }
}

static ERL_NIF_TERM mlx_eigvalsh(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    std::string uplo = "L";
    if (argc > 1) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        uplo = std::string(atom_name);
    }
    
    try {
        array result = linalg::eigvalsh(a, uplo);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "eigvalsh_error");
    }
}

// ==================== QR DECOMPOSITION ====================

static ERL_NIF_TERM mlx_qr(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    std::string mode = "reduced";
    if (argc > 1) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        mode = std::string(atom_name);
    }
    
    try {
        // MLX QR doesn't support mode parameter
        auto result = linalg::qr(a);
        std::vector<array> arrays = {std::get<0>(result), std::get<1>(result)};
        return make_array_tuple(env, arrays);
    } catch (const std::exception& e) {
        return make_error(env, "qr_error");
    }
}

// ==================== CHOLESKY DECOMPOSITION ====================

static ERL_NIF_TERM mlx_cholesky(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    bool upper = false;
    if (argc > 1) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        upper = (strcmp(atom_name, "true") == 0);
    }
    
    try {
        array result = linalg::cholesky(a, upper);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "cholesky_error");
    }
}

// ==================== MATRIX INVERSION ====================

static ERL_NIF_TERM mlx_inv(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = linalg::inv(a);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "inv_error");
    }
}

static ERL_NIF_TERM mlx_pinv(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 2) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    double rcond = 1e-15;
    if (argc > 1) {
        if (!enif_get_double(env, argv[1], &rcond)) {
            return enif_make_badarg(env);
        }
    }
    
    try {
        array result = linalg::pinv(a);  // MLX pinv doesn't take rcond parameter
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "pinv_error");
    }
}

// ==================== MATRIX NORMS ====================

static ERL_NIF_TERM mlx_norm(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 3) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    // Default parameters
    std::string ord = "fro";
    std::vector<int> axis;
    bool keepdims = false;
    
    if (argc > 1) {
        char atom_name[32];
        if (enif_get_atom(env, argv[1], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            ord = std::string(atom_name);
        } else {
            int ord_int;
            if (enif_get_int(env, argv[1], &ord_int)) {
                ord = std::to_string(ord_int);
            } else {
                return enif_make_badarg(env);
            }
        }
    }
    
    if (argc > 2) {
        // Parse axis - can be integer or list of integers
        int single_axis;
        if (enif_get_int(env, argv[2], &single_axis)) {
            axis = {single_axis};
        } else {
            unsigned int axis_len;
            if (enif_get_list_length(env, argv[2], &axis_len)) {
                axis.resize(axis_len);
                ERL_NIF_TERM head, tail = argv[2];
                for (unsigned int i = 0; i < axis_len; i++) {
                    if (!enif_get_list_cell(env, tail, &head, &tail) ||
                        !enif_get_int(env, head, &axis[i])) {
                        return enif_make_badarg(env);
                    }
                }
            } else {
                return enif_make_badarg(env);
            }
        }
    }
    
    try {
        array result = array({1.0f}); // Initialize with dummy value
        if (axis.empty()) {
            result = linalg::norm(a, ord, std::nullopt, keepdims);
        } else {
            result = linalg::norm(a, ord, axis, keepdims);
        }
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "norm_error");
    }
}

// ==================== MATRIX DETERMINANT ====================

static ERL_NIF_TERM mlx_det(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        // TODO: MLX may add det in future versions
        return make_error(env, "det_not_yet_available");
    } catch (const std::exception& e) {
        return make_error(env, "det_error");
    }
}

static ERL_NIF_TERM mlx_slogdet(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc != 1) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    try {
        // TODO: MLX may add slogdet in future versions
        return make_error(env, "slogdet_not_yet_available");
    } catch (const std::exception& e) {
        return make_error(env, "slogdet_error");
    }
}

// ==================== MATRIX TRACE ====================

static ERL_NIF_TERM mlx_trace(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 1 || argc > 4) return enif_make_badarg(env);
    
    array a = array({1.0f}); // Initialize with dummy value
    if (!get_array_resource(env, argv[0], a)) {
        return enif_make_badarg(env);
    }
    
    int offset = 0;
    int axis1 = -2;
    int axis2 = -1;
    
    if (argc > 1 && !enif_get_int(env, argv[1], &offset)) {
        return enif_make_badarg(env);
    }
    
    if (argc > 2 && !enif_get_int(env, argv[2], &axis1)) {
        return enif_make_badarg(env);
    }
    
    if (argc > 3 && !enif_get_int(env, argv[3], &axis2)) {
        return enif_make_badarg(env);
    }
    
    try {
        // TODO: MLX may add trace in future versions
        return make_error(env, "trace_not_yet_available");
    } catch (const std::exception& e) {
        return make_error(env, "trace_error");
    }
}

// ==================== CROSS PRODUCT ====================

static ERL_NIF_TERM mlx_cross(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 2 || argc > 3) return enif_make_badarg(env);
    
    array a = array({1.0f}), b = array({1.0f}); // Initialize with dummy values
    if (!get_array_resource(env, argv[0], a) ||
        !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    int axis = -1;
    if (argc > 2 && !enif_get_int(env, argv[2], &axis)) {
        return enif_make_badarg(env);
    }
    
    try {
        array result = linalg::cross(a, b, axis);
        return make_array_resource(env, result);
    } catch (const std::exception& e) {
        return make_error(env, "cross_error");
    }
}

// ==================== TRIANGULAR SOLVE ====================

static ERL_NIF_TERM mlx_tri_solve(ErlNifEnv* env, int argc, const ERL_NIF_TERM argv[]) {
    if (argc < 2 || argc > 4) return enif_make_badarg(env);
    
    array a = array({1.0f}), b = array({1.0f}); // Initialize with dummy values
    if (!get_array_resource(env, argv[0], a) ||
        !get_array_resource(env, argv[1], b)) {
        return enif_make_badarg(env);
    }
    
    bool upper = false;
    bool transpose = false;
    
    if (argc > 2) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[2], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        upper = (strcmp(atom_name, "true") == 0);
    }
    
    if (argc > 3) {
        char atom_name[32];
        if (!enif_get_atom(env, argv[3], atom_name, sizeof(atom_name), ERL_NIF_LATIN1)) {
            return enif_make_badarg(env);
        }
        transpose = (strcmp(atom_name, "true") == 0);
    }
    
    try {
        // TODO: MLX may add tri_solve in future versions
        return make_error(env, "tri_solve_not_yet_available");
    } catch (const std::exception& e) {
        return make_error(env, "tri_solve_error");
    }
}

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

// Linear algebra NIF function table
static ErlNifFunc nif_funcs[] = {
    // SVD
    {"svd", 1, mlx_svd, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"svd", 2, mlx_svd, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Eigenvalue decomposition
    {"eig", 1, mlx_eig, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"eigvals", 1, mlx_eigvals, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"eigh", 1, mlx_eigh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"eigh", 2, mlx_eigh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"eigvalsh", 1, mlx_eigvalsh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"eigvalsh", 2, mlx_eigvalsh, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // QR decomposition
    {"qr", 1, mlx_qr, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"qr", 2, mlx_qr, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Cholesky decomposition
    {"cholesky", 1, mlx_cholesky, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"cholesky", 2, mlx_cholesky, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Matrix inversion
    {"inv", 1, mlx_inv, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"pinv", 1, mlx_pinv, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"pinv", 2, mlx_pinv, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Matrix norms
    {"norm", 1, mlx_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"norm", 2, mlx_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"norm", 3, mlx_norm, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Matrix determinant
    {"det", 1, mlx_det, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"slogdet", 1, mlx_slogdet, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Matrix trace
    {"trace", 1, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"trace", 2, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"trace", 3, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"trace", 4, mlx_trace, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Cross product
    {"cross", 2, mlx_cross, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"cross", 3, mlx_cross, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    
    // Triangular solve
    {"tri_solve", 2, mlx_tri_solve, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"tri_solve", 3, mlx_tri_solve, ERL_NIF_DIRTY_JOB_CPU_BOUND},
    {"tri_solve", 4, mlx_tri_solve, 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_linalg_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_linalg_nif, nif_funcs, load, NULL, upgrade, unload)