#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