BareGit
#include "api_server.hpp"

#include <algorithm>
#include <chrono>
#include <cctype>
#include <charconv>
#include <cstdint>
#include <format>
#include <initializer_list>
#include <string>
#include <utility>

#include <mw/url.hpp>

#include "service_error.hpp"

namespace telegrammer
{

namespace
{

std::string lower(std::string value)
{
    std::transform(value.begin(), value.end(), value.begin(),
                   [](unsigned char character)
                   {
                       return static_cast<char>(std::tolower(character));
                   });
    return value;
}

std::optional<std::string_view> allowedMethods(std::string_view path)
{
    if(path == "/send")
    {
        return "POST";
    }
    if(path == "/subscribe")
    {
        return "POST";
    }
    if(path == "/subscriptions")
    {
        return "GET";
    }
    if(path == "/health")
    {
        return "GET";
    }
    if(path.starts_with("/subscriptions/") &&
       path.size() > std::string_view("/subscriptions/").size() &&
       path.find('/', std::string_view("/subscriptions/").size()) ==
           std::string_view::npos)
    {
        return "DELETE";
    }
    return std::nullopt;
}

httplib::TaskQueue* createRequestQueue()
{
    return new httplib::ThreadPool(8, 8, 64);
}

void configureServerSocket(socket_t socket)
{
#ifdef SO_REUSEPORT
    [[maybe_unused]] bool reuse_port =
        httplib::set_socket_opt(socket, SOL_SOCKET, SO_REUSEPORT, 0);
#else
    [[maybe_unused]] bool reuse_address =
        httplib::set_socket_opt(socket, SOL_SOCKET, SO_REUSEADDR, 0);
#endif
}

bool hasOnlyFields(const json& body,
                   std::initializer_list<std::string_view> fields)
{
    for(const auto& item: body.items())
    {
        bool known = false;
        for(std::string_view field: fields)
        {
            if(item.key() == field)
            {
                known = true;
                break;
            }
        }
        if(!known)
        {
            return false;
        }
    }
    return true;
}

bool getInt64(const json& body, const char* name, int64_t& result)
{
    if(!body.contains(name) || !body[name].is_number_integer())
    {
        return false;
    }
    try
    {
        result = body[name].get<int64_t>();
    }
    catch(const std::exception&)
    {
        return false;
    }
    return result != 0;
}

bool validUtf8(std::string_view value, std::size_t& code_points)
{
    code_points = 0;
    for(std::size_t i = 0; i < value.size(); ++code_points)
    {
        unsigned char first = static_cast<unsigned char>(value[i]);
        std::size_t length = 0;
        uint32_t code_point = 0;
        if(first <= 0x7f)
        {
            length = 1;
            code_point = first;
        }
        else if(first >= 0xc2 && first <= 0xdf)
        {
            length = 2;
            code_point = first & 0x1f;
        }
        else if(first >= 0xe0 && first <= 0xef)
        {
            length = 3;
            code_point = first & 0x0f;
        }
        else if(first >= 0xf0 && first <= 0xf4)
        {
            length = 4;
            code_point = first & 0x07;
        }
        else
        {
            return false;
        }
        if(i + length > value.size())
        {
            return false;
        }
        for(std::size_t j = 1; j < length; ++j)
        {
            unsigned char continuation =
                static_cast<unsigned char>(value[i + j]);
            if((continuation & 0xc0) != 0x80)
            {
                return false;
            }
            code_point = (code_point << 6) | (continuation & 0x3f);
        }
        if((length == 2 && code_point < 0x80) ||
           (length == 3 && code_point < 0x800) ||
           (length == 4 && code_point < 0x10000) ||
           code_point > 0x10ffff ||
           (code_point >= 0xd800 && code_point <= 0xdfff))
        {
            return false;
        }
        i += length;
    }
    return code_points <= 4096;
}

std::optional<std::string> normalizedUsername(const std::string& input)
{
    if(input.empty() || input.size() > 64)
    {
        return std::nullopt;
    }
    std::string result = input;
    if(result.starts_with('@'))
    {
        result.erase(0, 1);
    }
    if(result.empty() || result.size() > 64)
    {
        return std::nullopt;
    }
    for(char character: result)
    {
        bool valid = (character >= 'a' && character <= 'z') ||
                     (character >= 'A' && character <= 'Z') ||
                     (character >= '0' && character <= '9') ||
                     character == '_';
        if(!valid)
        {
            return std::nullopt;
        }
    }
    return lower(result);
}

mw::E<std::string> canonicalCallbackUrl(const std::string& input)
{
    if(input.empty() || input.size() > 2048)
    {
        return std::unexpected(serviceError(
            400, "INVALID_CALLBACK_URL", "Callback URL must be 1 to 2048 bytes",
            std::nullopt, false));
    }
    for(unsigned char character: input)
    {
        if(std::iscntrl(character) || std::isspace(character))
        {
            return std::unexpected(serviceError(
                400, "INVALID_CALLBACK_URL", "Callback URL contains whitespace"));
        }
    }
    auto parsed = mw::URL::fromStr(input);
    if(!parsed.has_value() || !parsed->valid())
    {
        return std::unexpected(serviceError(
            400, "INVALID_CALLBACK_URL", "Callback URL is invalid"));
    }
    std::string scheme = lower(parsed->scheme());
    if((scheme != "http" && scheme != "https") || parsed->host().empty() ||
       !parsed->user().empty() || !parsed->password().empty() ||
       !parsed->fragment().empty())
    {
        return std::unexpected(serviceError(
            400, "INVALID_CALLBACK_URL", "Callback URL must be HTTP or HTTPS"));
    }
    return parsed->str();
}

} // namespace

ApiServer::ApiServer(mw::IPSocketInfo listen_info, KeyStore& keys,
                     SubscriptionStore& subscriptions,
                     DeliveryStore& deliveries, TelegramApi& telegram,
                     RuntimeState& state)
    : mw::HTTPServer(listen_info), listen_info_(std::move(listen_info)),
      keys_(keys), subscriptions_(subscriptions), deliveries_(deliveries),
      telegram_(telegram), state_(state)
{}

ApiServer::~ApiServer()
{
    stop();
    wait();
}

mw::E<void> ApiServer::start()
{
    setup();
    if(!server.bind_to_port(listen_info_.address, listen_info_.port))
    {
        return std::unexpected(mw::runtimeError(
            std::format("Unable to bind {}:{}", listen_info_.address,
                        listen_info_.port)));
    }
    server_thread_ = std::thread([this]()
                                  {
                                      server.listen_after_bind();
                                  });
    server.wait_until_ready();
    return {};
}

void ApiServer::stop()
{
    server.stop();
}

void ApiServer::wait()
{
    if(server_thread_.joinable())
    {
        server_thread_.join();
    }
}

void ApiServer::setup()
{
    server.new_task_queue = createRequestQueue;
    server.set_socket_options(configureServerSocket);
    server.set_payload_max_length(64 * 1024);
    server.set_read_timeout(std::chrono::seconds(5));
    server.set_write_timeout(std::chrono::seconds(5));
    server.set_default_headers({{"Cache-Control", "no-store"}});
    server.set_error_handler([]([[maybe_unused]] const Request& request,
                                Response& response)
                             {
                                 if(!response.body.empty())
                                 {
                                     return;
                                 }
                                 int status = response.status;
                                 writeError(response, status,
                                            status == 404 ? "NOT_FOUND" :
                                            status == 405 ? "METHOD_NOT_ALLOWED" :
                                            status == 413 ? "PAYLOAD_TOO_LARGE" :
                                            "HTTP_ERROR",
                                            status == 404 ? "Route not found" :
                                            status == 405 ? "Method not allowed" :
                                            status == 413 ? "Request body is too large" :
                                            "HTTP request failed");
                             });
    server.set_exception_handler(
        [](const Request& request, Response& response,
           std::exception_ptr exception)
        {
            [[maybe_unused]] const Request& ignored_request = request;
            [[maybe_unused]] std::exception_ptr ignored_exception = exception;
            writeError(response, 500, "INTERNAL_ERROR",
                       "The request could not be completed");
        });
    server.set_pre_request_handler(
        [this](const Request& request, Response& response)
        {
            if(!authenticate(request, response))
            {
                return httplib::Server::HandlerResponse::Handled;
            }
            return httplib::Server::HandlerResponse::Unhandled;
        });
    server.set_pre_routing_handler(
        [this](const Request& request, Response& response)
        {
            auto methods = allowedMethods(request.path);
            if(!methods.has_value() || request.method == *methods)
            {
                return httplib::Server::HandlerResponse::Unhandled;
            }
            if(!authenticate(request, response))
            {
                return httplib::Server::HandlerResponse::Handled;
            }
            response.status = 405;
            response.set_header("Allow", std::string(*methods));
            return httplib::Server::HandlerResponse::Handled;
        });

    server.Post("/send", [this](const Request& request, Response& response)
                {
                    send(request, response);
                });
    server.Post("/subscribe",
                [this](const Request& request, Response& response)
                {
                    subscribe(request, response);
                });
    server.Get("/subscriptions",
               [this](const Request& request, Response& response)
               {
                   listSubscriptions(request, response);
               });
    server.Delete("/subscriptions/:id",
                  [this](const Request& request, Response& response)
                  {
                      deleteSubscription(request, response);
                  });
    server.Get("/health",
               [this](const Request& request, Response& response)
               {
                   health(request, response);
               });
}

bool ApiServer::authenticate(const Request& request, Response& response)
{
    if(request.get_header_value_count("Authorization") != 1)
    {
        writeError(response, 401, "UNAUTHORIZED", "Bearer credentials required");
        response.set_header("WWW-Authenticate", "Bearer");
        return false;
    }
    std::string authorization = request.get_header_value("Authorization");
    if(authorization.size() <= 7 ||
       lower(authorization.substr(0, 7)) != "bearer " ||
       authorization.find_first_of(" \t\r\n", 7) != std::string::npos)
    {
        writeError(response, 401, "UNAUTHORIZED", "Bearer credentials required");
        response.set_header("WWW-Authenticate", "Bearer");
        return false;
    }

    auto identity = keys_.authenticate(authorization.substr(7));
    if(!identity.has_value())
    {
        writeError(response, identity.error());
        return false;
    }
    if(!identity->has_value())
    {
        writeError(response, 401, "UNAUTHORIZED", "Invalid bearer credential");
        response.set_header("WWW-Authenticate", "Bearer");
        return false;
    }
    response.user_data.set("key_id", identity->value().id);
    response.user_data.set("key_name", identity->value().name);
    return true;
}

bool ApiServer::requireJson(const Request& request, Response& response)
{
    if(!request.has_header("Content-Type"))
    {
        writeError(response, 415, "UNSUPPORTED_MEDIA_TYPE",
                   "Content-Type must be application/json");
        return false;
    }
    std::string content_type = lower(request.get_header_value("Content-Type"));
    std::size_t separator = content_type.find(';');
    std::string media_type = content_type.substr(0, separator);
    while(!media_type.empty() && media_type.back() == ' ')
    {
        media_type.pop_back();
    }
    if(media_type != "application/json")
    {
        writeError(response, 415, "UNSUPPORTED_MEDIA_TYPE",
                   "Content-Type must be application/json");
        return false;
    }
    return true;
}

std::optional<json> ApiServer::parseObject(const Request& request,
                                            Response& response)
{
    if(!requireJson(request, response))
    {
        return std::nullopt;
    }
    try
    {
        json body = json::parse(request.body);
        if(!body.is_object())
        {
            writeError(response, 400, "INVALID_REQUEST",
                       "Request body must be a JSON object");
            return std::nullopt;
        }
        return body;
    }
    catch(const std::exception&)
    {
        writeError(response, 400, "INVALID_REQUEST", "Request body is invalid JSON");
        return std::nullopt;
    }
}

std::optional<int64_t> ApiServer::ownerId(const Response& response) const
{
    const int64_t* value = response.user_data.get<int64_t>("key_id");
    if(value == nullptr)
    {
        return std::nullopt;
    }
    return *value;
}

void ApiServer::writeSuccess(Response& response, const json& body, int status)
{
    response.status = status;
    response.set_content(body.dump(), "application/json");
    response.set_header("Cache-Control", "no-store");
}

void ApiServer::writeNoContent(Response& response)
{
    response.status = 204;
    response.set_header("Cache-Control", "no-store");
}

void ApiServer::writeError(Response& response, int status,
                           std::string_view code, std::string_view message,
                           std::optional<std::string_view> field,
                           std::optional<int> retry_after)
{
    json error = {{"ok", false}};
    error["error"] = {{"code", code}, {"message", message}};
    if(field.has_value())
    {
        error["error"]["field"] = *field;
    }
    response.status = status;
    response.set_content(error.dump(), "application/json");
    response.set_header("Cache-Control", "no-store");
    if(retry_after.has_value())
    {
        response.set_header("Retry-After", std::to_string(*retry_after));
    }
}

void ApiServer::writeError(Response& response, const mw::Error& error)
{
    if(const ServiceError* service_error = asServiceError(error);
       service_error != nullptr)
    {
        writeError(response, service_error->status, service_error->code,
                   service_error->msg, std::nullopt,
                   service_error->retry_after);
        return;
    }
    writeError(response, 500, "INTERNAL_ERROR",
               "The request could not be completed");
}

void ApiServer::send(const Request& request, Response& response)
{
    auto body = parseObject(request, response);
    if(!body.has_value())
    {
        return;
    }
    if(!hasOnlyFields(*body, {"chat_id", "username", "text"}) ||
       !body->contains("text") || !(*body)["text"].is_string())
    {
        writeError(response, 400, "INVALID_REQUEST",
                   "Request must contain only chat_id, username, and text");
        return;
    }
    std::string text = (*body)["text"].get<std::string>();
    std::size_t code_points = 0;
    if(text.empty() || !validUtf8(text, code_points))
    {
        writeError(response, 400, "INVALID_REQUEST",
                   "text must be valid nonempty UTF-8", "text");
        return;
    }

    bool has_chat_id = body->contains("chat_id");
    bool has_username = body->contains("username");
    if(has_chat_id == has_username)
    {
        writeError(response, 400, "INVALID_REQUEST",
                   "Exactly one of chat_id and username is required");
        return;
    }

    int64_t chat_id = 0;
    if(has_chat_id)
    {
        if(!getInt64(*body, "chat_id", chat_id))
        {
            writeError(response, 400, "INVALID_REQUEST",
                       "chat_id must be a nonzero signed integer", "chat_id");
            return;
        }
    }
    else
    {
        if(!(*body)["username"].is_string())
        {
            writeError(response, 400, "INVALID_REQUEST",
                       "username must be a string", "username");
            return;
        }
        auto username = normalizedUsername(
            (*body)["username"].get<std::string>());
        if(!username.has_value())
        {
            writeError(response, 400, "INVALID_REQUEST",
                       "username is invalid", "username");
            return;
        }
        auto resolved = deliveries_.resolveUsername(*username);
        if(!resolved.has_value())
        {
            writeError(response, resolved.error());
            return;
        }
        if(!resolved->has_value())
        {
            writeError(response, 404, "USERNAME_NOT_FOUND",
                       "Username has not been observed in a private chat");
            return;
        }
        chat_id = **resolved;
    }

    auto result = telegram_.sendMessage(chat_id, text);
    if(!result.has_value())
    {
        writeError(response, result.error());
        return;
    }
    writeSuccess(response, *result);
}

void ApiServer::subscribe(const Request& request, Response& response)
{
    auto body = parseObject(request, response);
    if(!body.has_value())
    {
        return;
    }
    if(!hasOnlyFields(*body, {"chat_id", "callback_url"}) ||
       !body->contains("chat_id") || !body->contains("callback_url"))
    {
        writeError(response, 400, "INVALID_REQUEST",
                   "Request must contain chat_id and callback_url");
        return;
    }
    int64_t chat_id = 0;
    if(!getInt64(*body, "chat_id", chat_id))
    {
        writeError(response, 400, "INVALID_REQUEST",
                   "chat_id must be a nonzero signed integer", "chat_id");
        return;
    }
    if(!(*body)["callback_url"].is_string())
    {
        writeError(response, 400, "INVALID_REQUEST",
                   "callback_url must be a string", "callback_url");
        return;
    }
    auto url = canonicalCallbackUrl(
        (*body)["callback_url"].get<std::string>());
    if(!url.has_value())
    {
        writeError(response, url.error());
        return;
    }
    auto owner = ownerId(response);
    if(!owner.has_value())
    {
        writeError(response, 500, "INTERNAL_ERROR", "Missing request identity");
        return;
    }
    auto result = subscriptions_.add(*owner, chat_id, *url);
    if(!result.has_value())
    {
        writeError(response, result.error());
        return;
    }
    writeSuccess(response, { {"ok", true}, {"subscription_id", result->id} });
}

void ApiServer::listSubscriptions(const Request& request, Response& response)
{
    [[maybe_unused]] const Request& ignored_request = request;
    auto owner = ownerId(response);
    if(!owner.has_value())
    {
        writeError(response, 500, "INTERNAL_ERROR", "Missing request identity");
        return;
    }
    auto result = subscriptions_.list(*owner);
    if(!result.has_value())
    {
        writeError(response, result.error());
        return;
    }
    json subscriptions = json::array();
    for(const Subscription& subscription: *result)
    {
        subscriptions.push_back({
            {"id", subscription.id}, {"chat_id", subscription.chat_id},
            {"callback_url", subscription.callback_url},
            {"created_at", subscription.created_at}});
    }
    writeSuccess(response,
                 {{"ok", true}, {"subscriptions", subscriptions}});
}

void ApiServer::deleteSubscription(const Request& request,
                                   Response& response)
{
    auto owner = ownerId(response);
    if(!owner.has_value())
    {
        writeError(response, 500, "INTERNAL_ERROR", "Missing request identity");
        return;
    }
    auto iterator = request.path_params.find("id");
    int64_t id = 0;
    if(iterator == request.path_params.end())
    {
        writeError(response, 400, "INVALID_REQUEST", "Subscription ID is invalid");
        return;
    }
    const std::string& value = iterator->second;
    const char* end = value.data() + value.size();
    auto parsed = std::from_chars(value.data(), end, id);
    if(parsed.ec != std::errc{} || parsed.ptr != end || id <= 0)
    {
        writeError(response, 400, "INVALID_REQUEST", "Subscription ID is invalid");
        return;
    }
    auto result = subscriptions_.remove(*owner, id);
    if(!result.has_value())
    {
        writeError(response, result.error());
        return;
    }
    if(!*result)
    {
        writeError(response, 404, "NOT_FOUND", "Subscription not found");
        return;
    }
    writeNoContent(response);
}

void ApiServer::health(const Request& request, Response& response)
{
    [[maybe_unused]] const Request& ignored_request = request;
    auto queue = deliveries_.count();
    if(!queue.has_value())
    {
        writeError(response, queue.error());
        return;
    }
    bool ready = state_.polling_ready.load();
    bool saturated = *queue >= 100000;
    bool degraded = state_.degraded.load() || saturated;
    json body = {{"ok", ready && !degraded},
                 {"polling_ready", ready},
                 {"degraded", degraded},
                 {"queue_saturated", saturated},
                 {"queue_size", *queue},
                 {"last_success", state_.last_success.load()}};
    writeSuccess(response, body, ready && !degraded ? 200 : 503);
}

} // namespace telegrammer