BareGit
#include "database.hpp"

#include <array>
#include <cerrno>
#include <chrono>
#include <cstring>
#include <fcntl.h>
#include <optional>
#include <set>
#include <stdexcept>
#include <string_view>
#include <sys/file.h>
#include <sys/stat.h>
#include <unistd.h>
#include <vector>

#include "service_error.hpp"

namespace telegrammer
{

namespace
{

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

enum class TableShape
{
    MISSING,
    CURRENT,
    LEGACY,
    INCOMPATIBLE
};

TableShape tableShape(mw::SQLite& db, const char* table,
                      const std::set<std::string>& expected_columns)
{
    auto table_rows = db.eval<std::string>(
        std::string("SELECT name FROM sqlite_master WHERE type = 'table' "
                    "AND name = '") + table + "';");
    if(!table_rows.has_value())
    {
        return TableShape::INCOMPATIBLE;
    }

    // The table names passed here are fixed internal values. SQLite does not
    // accept a bound parameter in PRAGMA table_info(), so keep the values
    // out of user input and quote them explicitly.
    auto columns = db.eval<int64_t, std::string, std::string, int64_t,
                           std::optional<std::string>, int64_t>(
        std::string("PRAGMA table_info('") + table + "');");
    if(!columns.has_value())
    {
        return TableShape::INCOMPATIBLE;
    }
    if(columns->empty())
    {
        return TableShape::MISSING;
    }

    std::set<std::string> actual_columns;
    for(const auto& column: *columns)
    {
        actual_columns.insert(std::get<1>(column));
    }
    if(std::strcmp(table, "api_keys") == 0 &&
       actual_columns.contains("key") &&
       !actual_columns.contains("key_digest"))
    {
        return TableShape::LEGACY;
    }
    return actual_columns == expected_columns ? TableShape::CURRENT
                                               : TableShape::INCOMPATIBLE;
}

} // namespace

DatabaseLock::DatabaseLock(const std::string& db_path)
    : lock_path_(db_path + ".lock")
{
    fd_ = ::open(lock_path_.c_str(),
                 O_CREAT | O_RDWR | O_CLOEXEC | O_NOFOLLOW,
                 0600);
    if(fd_ < 0)
    {
        throw std::runtime_error(
            "Failed to open database lock '" + lock_path_ + "': " +
            std::strerror(errno));
    }
    if(::fchmod(fd_, 0600) != 0)
    {
        int error_code = errno;
        ::close(fd_);
        fd_ = -1;
        throw std::runtime_error(
            "Failed to secure database lock '" + lock_path_ + "': " +
            std::strerror(error_code));
    }

    if(::flock(fd_, LOCK_EX | LOCK_NB) == 0)
    {
        return;
    }

    int error_code = errno;
    ::close(fd_);
    fd_ = -1;

    if(error_code == EWOULDBLOCK || error_code == EAGAIN)
    {
        throw std::runtime_error(
            "Another Telegrammer daemon is using database '" + db_path +
            "'");
    }

    throw std::runtime_error(
        "Failed to lock database '" + db_path + "': " +
        std::strerror(error_code));
}

DatabaseLock::~DatabaseLock()
{
    if(fd_ >= 0)
    {
        ::close(fd_);
    }
}

mw::E<std::unique_ptr<mw::SQLite>> openDatabase(const std::string& db_path)
{
    auto result = mw::SQLite::connectFile(db_path, 5000);
    if(!result.has_value())
    {
        return std::unexpected(serviceError(
            503, "DATABASE_UNAVAILABLE", "The state database is unavailable"));
    }

    auto db = std::move(*result);
    for(const char* pragma: {
            "PRAGMA foreign_keys = ON;",
            "PRAGMA synchronous = NORMAL;"})
    {
        auto pragma_result = db->execute(pragma);
        if(!pragma_result.has_value())
        {
            return std::unexpected(serviceError(
                503, "DATABASE_UNAVAILABLE",
                "The state database is unavailable"));
        }
    }

    return db;
}

mw::E<void> initializeDatabase(const std::string& db_path)
{
    auto result = openDatabase(db_path);
    if(!result.has_value())
    {
        return std::unexpected(result.error());
    }
    mw::SQLite& db = **result;

    auto user_version = db.evalToValue<int64_t>("PRAGMA user_version;");
    if(!user_version.has_value() || *user_version != 0)
    {
        return std::unexpected(serviceError(
            500, "DATABASE_SCHEMA",
            "The database has an unsupported schema version."));
    }

    const std::set<std::string> api_keys_columns = {
        "id", "name", "key_digest", "created_at"};
    const std::set<std::string> subscriptions_columns = {
        "id", "owner_id", "chat_id", "callback_url", "created_at"};
    const std::set<std::string> poll_state_columns = {
        "singleton", "bot_id", "next_offset"};
    const std::set<std::string> usernames_columns = {
        "username", "chat_id", "observed_at"};
    const std::set<std::string> deliveries_columns = {
        "id", "subscription_id", "update_id", "payload", "state",
        "attempt_count", "next_attempt_at", "created_at", "last_error"};

    const TableShape api_keys_shape =
        tableShape(db, "api_keys", api_keys_columns);
    if(api_keys_shape == TableShape::LEGACY)
    {
        return std::unexpected(serviceError(
            500, "DATABASE_SCHEMA",
            "The database uses the development API-key schema. "
            "Remove it before starting the new development build."));
    }
    if(api_keys_shape == TableShape::INCOMPATIBLE)
    {
        return std::unexpected(serviceError(
            500, "DATABASE_SCHEMA",
            "The database has an incompatible api_keys table."));
    }

    const std::array<std::pair<const char*, std::set<std::string>>, 4>
        existing_tables = {{{"subscriptions", subscriptions_columns},
                            {"poll_state", poll_state_columns},
                            {"usernames", usernames_columns},
                            {"deliveries", deliveries_columns}}};
    for(const auto& [table, columns]: existing_tables)
    {
        if(tableShape(db, table, columns) == TableShape::INCOMPATIBLE)
        {
            return std::unexpected(serviceError(
                500, "DATABASE_SCHEMA",
                std::string("The database has an incompatible ") + table +
                    " table."));
        }
    }

    auto begin_result = execute(db, "BEGIN IMMEDIATE;");
    if(!begin_result.has_value())
    {
        return std::unexpected(serviceError(
            503, "DATABASE_UNAVAILABLE", "Unable to initialize the database"));
    }
    auto rollback = [&db]()
    {
        [[maybe_unused]] auto result = db.execute("ROLLBACK;");
    };

    // SQLite's prepare API executes one statement per call. Keep the schema
    // statements separate so every failure can be reported precisely.
    const std::vector<std::string> statements = {
        "CREATE TABLE IF NOT EXISTS api_keys ("
        "id INTEGER PRIMARY KEY, "
        "name TEXT NOT NULL UNIQUE, "
        "key_digest TEXT NOT NULL UNIQUE, "
        "created_at INTEGER NOT NULL"
        ");",
        "CREATE TABLE IF NOT EXISTS subscriptions ("
        "id INTEGER PRIMARY KEY, "
        "owner_id INTEGER NOT NULL REFERENCES api_keys(id) "
        "ON DELETE CASCADE, "
        "chat_id INTEGER NOT NULL, "
        "callback_url TEXT NOT NULL, "
        "created_at INTEGER NOT NULL, "
        "UNIQUE(owner_id, chat_id, callback_url)"
        ");",
        "CREATE TABLE IF NOT EXISTS poll_state ("
        "singleton INTEGER PRIMARY KEY CHECK(singleton = 1), "
        "bot_id INTEGER NOT NULL, "
        "next_offset INTEGER NOT NULL"
        ");",
        "CREATE TABLE IF NOT EXISTS usernames ("
        "username TEXT PRIMARY KEY, "
        "chat_id INTEGER NOT NULL UNIQUE, "
        "observed_at INTEGER NOT NULL"
        ");",
        "CREATE TABLE IF NOT EXISTS deliveries ("
        "id INTEGER PRIMARY KEY, "
        "subscription_id INTEGER NOT NULL "
        "REFERENCES subscriptions(id) ON DELETE CASCADE, "
        "update_id INTEGER NOT NULL, "
        "payload TEXT NOT NULL, "
        "state TEXT NOT NULL CHECK(state IN ('PENDING','IN_FLIGHT','DEAD')), "
        "attempt_count INTEGER NOT NULL DEFAULT 0, "
        "next_attempt_at INTEGER NOT NULL, "
        "created_at INTEGER NOT NULL, "
        "last_error TEXT, "
        "UNIQUE(subscription_id, update_id)"
        ");",
        "CREATE INDEX IF NOT EXISTS deliveries_due "
        "ON deliveries(state, next_attempt_at, id);"};

    for(const std::string& statement: statements)
    {
        auto statement_result = execute(db, statement);
        if(!statement_result.has_value())
        {
            rollback();
            return std::unexpected(serviceError(
                503, "DATABASE_UNAVAILABLE", "Unable to initialize the database"));
        }
    }

    auto commit_result = execute(db, "COMMIT;");
    if(!commit_result.has_value())
    {
        rollback();
        return std::unexpected(serviceError(
            503, "DATABASE_UNAVAILABLE", "Unable to initialize the database"));
    }

    return {};
}

int64_t nowSeconds()
{
    return std::chrono::duration_cast<std::chrono::seconds>(
               std::chrono::system_clock::now().time_since_epoch())
        .count();
}

} // namespace telegrammer