#include "swift_binding.h"

#include <new>
#include <utility>

#include "zstd_codec.h"

struct uasm_zstd_result
{
    uasm::zstd::Result value;
};

struct uasm_zstd_compressor
{
    uasm_zstd_compressor(int level, int workers) : value(level, workers) {}
    uasm::zstd::Compressor value;
};

struct uasm_zstd_decompressor
{
    explicit uasm_zstd_decompressor(size_t maximumOutputSize)
        : value(maximumOutputSize) {}
    uasm::zstd::Decompressor value;
};

namespace
{
    uasm::zstd::Result InvalidArgument(const char *message)
    {
        uasm::zstd::Result result;
        result.errorKind = uasm::zstd::ErrorKind::kInvalidArgument;
        result.errorMessage = message;
        return result;
    }

    uasm_zstd_result_t *Wrap(uasm::zstd::Result result)
    {
        return new (std::nothrow) uasm_zstd_result_t{std::move(result)};
    }
} // namespace

extern "C" uasm_zstd_result_t *uasm_zstd_compress(
    const uint8_t *input,
    size_t inputSize,
    int32_t level)
{
    if (input == nullptr && inputSize != 0)
    {
        return Wrap(InvalidArgument("zstd compression input must not be null"));
    }
    return Wrap(uasm::zstd::Compress(input, inputSize, static_cast<int>(level)));
}

extern "C" uasm_zstd_result_t *uasm_zstd_compress_with_workers(
    const uint8_t *input,
    size_t inputSize,
    int32_t level,
    int32_t workers)
{
    if (input == nullptr && inputSize != 0)
    {
        return Wrap(InvalidArgument("zstd compression input must not be null"));
    }
    return Wrap(uasm::zstd::Compress(
        input, inputSize, static_cast<int>(level), static_cast<int>(workers)));
}

extern "C" uasm_zstd_result_t *uasm_zstd_decompress(
    const uint8_t *input,
    size_t inputSize,
    size_t maxOutputSize)
{
    if (input == nullptr && inputSize != 0)
    {
        return Wrap(InvalidArgument("zstd decompression input must not be null"));
    }
    return Wrap(uasm::zstd::Decompress(input, inputSize, maxOutputSize));
}

extern "C" uasm_zstd_compressor_t *uasm_zstd_compressor_create(int32_t level)
{
    return new (std::nothrow) uasm_zstd_compressor_t(static_cast<int>(level), 0);
}

extern "C" uasm_zstd_compressor_t *uasm_zstd_compressor_create_with_workers(
    int32_t level,
    int32_t workers)
{
    return new (std::nothrow) uasm_zstd_compressor_t(
        static_cast<int>(level), static_cast<int>(workers));
}

extern "C" uasm_zstd_result_t *uasm_zstd_compressor_status(
    const uasm_zstd_compressor_t *compressor)
{
    return compressor == nullptr
        ? Wrap(InvalidArgument("zstd compressor must not be null"))
        : Wrap(compressor->value.status());
}

extern "C" uasm_zstd_result_t *uasm_zstd_compressor_update(
    uasm_zstd_compressor_t *compressor,
    const uint8_t *input,
    size_t inputSize)
{
    if (compressor == nullptr)
    {
        return Wrap(InvalidArgument("zstd compressor must not be null"));
    }
    if (input == nullptr && inputSize != 0)
    {
        return Wrap(InvalidArgument("zstd compression input must not be null"));
    }
    return Wrap(compressor->value.Update(input, inputSize));
}

extern "C" uasm_zstd_result_t *uasm_zstd_compressor_finish(
    uasm_zstd_compressor_t *compressor)
{
    return compressor == nullptr
        ? Wrap(InvalidArgument("zstd compressor must not be null"))
        : Wrap(compressor->value.Finish());
}

extern "C" void uasm_zstd_compressor_destroy(uasm_zstd_compressor_t *compressor)
{
    delete compressor;
}

extern "C" uasm_zstd_decompressor_t *uasm_zstd_decompressor_create(
    size_t maxOutputSize)
{
    return new (std::nothrow) uasm_zstd_decompressor_t(maxOutputSize);
}

extern "C" uasm_zstd_result_t *uasm_zstd_decompressor_status(
    const uasm_zstd_decompressor_t *decompressor)
{
    return decompressor == nullptr
        ? Wrap(InvalidArgument("zstd decompressor must not be null"))
        : Wrap(decompressor->value.status());
}

extern "C" uasm_zstd_result_t *uasm_zstd_decompressor_update(
    uasm_zstd_decompressor_t *decompressor,
    const uint8_t *input,
    size_t inputSize)
{
    if (decompressor == nullptr)
    {
        return Wrap(InvalidArgument("zstd decompressor must not be null"));
    }
    if (input == nullptr && inputSize != 0)
    {
        return Wrap(InvalidArgument("zstd decompression input must not be null"));
    }
    return Wrap(decompressor->value.Update(input, inputSize));
}

extern "C" uasm_zstd_result_t *uasm_zstd_decompressor_finish(
    uasm_zstd_decompressor_t *decompressor)
{
    return decompressor == nullptr
        ? Wrap(InvalidArgument("zstd decompressor must not be null"))
        : Wrap(decompressor->value.Finish());
}

extern "C" void uasm_zstd_decompressor_destroy(uasm_zstd_decompressor_t *decompressor)
{
    delete decompressor;
}

extern "C" uasm_zstd_error_kind_t uasm_zstd_result_error_kind(
    const uasm_zstd_result_t *result)
{
    if (result == nullptr)
    {
        return UASM_ZSTD_ERROR_RUNTIME;
    }
    switch (result->value.errorKind)
    {
    case uasm::zstd::ErrorKind::kNone:
        return UASM_ZSTD_ERROR_NONE;
    case uasm::zstd::ErrorKind::kInvalidArgument:
        return UASM_ZSTD_ERROR_INVALID_ARGUMENT;
    case uasm::zstd::ErrorKind::kRange:
        return UASM_ZSTD_ERROR_RANGE;
    case uasm::zstd::ErrorKind::kRuntime:
        return UASM_ZSTD_ERROR_RUNTIME;
    }
    return UASM_ZSTD_ERROR_RUNTIME;
}

extern "C" const char *uasm_zstd_result_error_message(
    const uasm_zstd_result_t *result)
{
    if (result == nullptr)
    {
        return "unable to allocate zstd result";
    }
    return result->value.errorMessage.c_str();
}

extern "C" const uint8_t *uasm_zstd_result_data(
    const uasm_zstd_result_t *result)
{
    if (result == nullptr || !result->value || result->value.bytes.empty())
    {
        return nullptr;
    }
    return result->value.bytes.data();
}

extern "C" size_t uasm_zstd_result_size(const uasm_zstd_result_t *result)
{
    if (result == nullptr || !result->value)
    {
        return 0;
    }
    return result->value.bytes.size();
}

extern "C" void uasm_zstd_result_destroy(uasm_zstd_result_t *result)
{
    delete result;
}