BareGit
#include "key_store.hpp"

#include <array>
#include <format>
#include <openssl/evp.h>
#include <openssl/rand.h>

#include "database.hpp"
#include "service_error.hpp"

namespace telegrammer
{

namespace
{

mw::E<void> execute(mw::SQLite& db, const std::string& sql)
{
    auto result = db.execute(sql);
    if(!result.has_value())
    {
        return std::unexpected(result.error());
    }
    return {};
}

mw::Error databaseError([[maybe_unused]] const mw::Error& error)
{
    return serviceError(503, "DATABASE_UNAVAILABLE",
                        "The state database is unavailable");
}

bool isUniqueError(const mw::Error& error)
{
    return mw::errorMsg(error).find("UNIQUE") != std::string::npos;
}

} // namespace

KeyStore::KeyStore(std::string db_path)
    : db_path_(std::move(db_path))
{}

mw::E<std::string> KeyStore::generateKey()
{
    std::array<unsigned char, 32> bytes{};
    if(RAND_priv_bytes(bytes.data(), static_cast<int>(bytes.size())) != 1)
    {
        return std::unexpected(serviceError(
            500, "KEY_GENERATION_FAILED", "Unable to generate API key"));
    }

    std::string key;
    key.reserve(bytes.size() * 2);
    for(unsigned char byte: bytes)
    {
        key += std::format("{:02x}", byte);
    }
    return key;
}

mw::E<std::string> KeyStore::digestKey(const std::string& key)
{
    std::array<unsigned char, EVP_MAX_MD_SIZE> digest{};
    unsigned int digest_size = 0;
    EVP_MD_CTX* context = EVP_MD_CTX_new();
    if(context == nullptr)
    {
        return std::unexpected(serviceError(
            500, "KEY_HASH_FAILED", "Unable to initialize key hashing"));
    }

    bool success = EVP_DigestInit_ex(context, EVP_sha256(), nullptr) == 1 &&
                   EVP_DigestUpdate(context, key.data(), key.size()) == 1 &&
                   EVP_DigestFinal_ex(context, digest.data(), &digest_size) ==
                       1;
    EVP_MD_CTX_free(context);
    if(!success || digest_size != 32)
    {
        return std::unexpected(serviceError(
            500, "KEY_HASH_FAILED", "Unable to hash API key"));
    }

    std::string result;
    result.reserve(digest_size * 2);
    for(unsigned int i = 0; i < digest_size; ++i)
    {
        result += std::format("{:02x}", digest[i]);
    }
    return result;
}

mw::E<std::string> KeyStore::addKey(const std::string& name) const
{
    if(name.empty() || name.size() > 128)
    {
        return std::unexpected(serviceError(
            400, "INVALID_KEY_NAME", "Key name must be 1 to 128 bytes"));
    }

    auto key = generateKey();
    if(!key.has_value())
    {
        return std::unexpected(key.error());
    }
    auto digest = digestKey(*key);
    if(!digest.has_value())
    {
        return std::unexpected(digest.error());
    }

    auto db_result = openDatabase(db_path_);
    if(!db_result.has_value())
    {
        return std::unexpected(db_result.error());
    }
    mw::SQLite& db = **db_result;

    auto statement = db.statementFromStr(
        "INSERT INTO api_keys (name, key_digest, created_at) "
        "VALUES (?, ?, ?);");
    if(!statement.has_value())
    {
        return std::unexpected(databaseError(statement.error()));
    }
    auto bind_result = statement->bind(name, *digest, nowSeconds());
    if(!bind_result.has_value())
    {
        return std::unexpected(databaseError(bind_result.error()));
    }
    auto insert_result = db.execute(std::move(*statement));
    if(!insert_result.has_value())
    {
        if(isUniqueError(insert_result.error()))
        {
            return std::unexpected(serviceError(
                409, "KEY_EXISTS", "A key with that name already exists"));
        }
        return std::unexpected(databaseError(insert_result.error()));
    }

    return *key;
}

mw::E<std::optional<KeyIdentity>> KeyStore::authenticate(
    const std::string& key) const
{
    auto digest = digestKey(key);
    if(!digest.has_value())
    {
        return std::unexpected(digest.error());
    }

    auto db_result = openDatabase(db_path_);
    if(!db_result.has_value())
    {
        return std::unexpected(db_result.error());
    }
    mw::SQLite& db = **db_result;

    auto statement = db.statementFromStr(
        "SELECT id, name FROM api_keys WHERE key_digest = ? LIMIT 1;");
    if(!statement.has_value())
    {
        return std::unexpected(databaseError(statement.error()));
    }
    auto bind_result = statement->bind(*digest);
    if(!bind_result.has_value())
    {
        return std::unexpected(databaseError(bind_result.error()));
    }
    auto rows = db.eval<int64_t, std::string>(std::move(*statement));
    if(!rows.has_value())
    {
        return std::unexpected(databaseError(rows.error()));
    }
    if(rows->empty())
    {
        return std::nullopt;
    }
    return KeyIdentity{std::get<0>((*rows)[0]), std::get<1>((*rows)[0])};
}

mw::E<std::vector<KeyInfo>> KeyStore::listKeys() const
{
    auto db_result = openDatabase(db_path_);
    if(!db_result.has_value())
    {
        return std::unexpected(db_result.error());
    }
    auto rows = (*db_result)->eval<int64_t, std::string, int64_t>(
        "SELECT id, name, created_at FROM api_keys ORDER BY id;");
    if(!rows.has_value())
    {
        return std::unexpected(databaseError(rows.error()));
    }

    std::vector<KeyInfo> result;
    result.reserve(rows->size());
    for(const auto& row: *rows)
    {
        result.push_back(
            {std::get<0>(row), std::get<1>(row), std::get<2>(row)});
    }
    return result;
}

mw::E<bool> KeyStore::deleteKey(const std::string& name) const
{
    auto db_result = openDatabase(db_path_);
    if(!db_result.has_value())
    {
        return std::unexpected(db_result.error());
    }
    mw::SQLite& db = **db_result;

    auto begin_result = execute(db, "BEGIN IMMEDIATE;");
    if(!begin_result.has_value())
    {
        return std::unexpected(databaseError(begin_result.error()));
    }

    auto rollback = [&db]()
    {
        [[maybe_unused]] auto result = db.execute("ROLLBACK;");
    };

    auto statement = db.statementFromStr(
        "DELETE FROM api_keys WHERE name = ?;");
    if(!statement.has_value())
    {
        rollback();
        return std::unexpected(databaseError(statement.error()));
    }
    auto bind_result = statement->bind(name);
    if(!bind_result.has_value())
    {
        rollback();
        return std::unexpected(databaseError(bind_result.error()));
    }
    auto delete_result = db.execute(std::move(*statement));
    if(!delete_result.has_value())
    {
        rollback();
        return std::unexpected(databaseError(delete_result.error()));
    }
    bool deleted = db.changedRowsCount() > 0;

    auto commit_result = execute(db, "COMMIT;");
    if(!commit_result.has_value())
    {
        rollback();
        return std::unexpected(databaseError(commit_result.error()));
    }
    return deleted;
}

} // namespace telegrammer