#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