#include "delivery_store.hpp"
#include <algorithm>
#include <chrono>
#include <limits>
#include <random>
#include <string_view>
#include <tuple>
#include <utility>
#include "database.hpp"
#include "service_error.hpp"
namespace telegrammer
{
namespace
{
struct ValidatedUpdate
{
int64_t update_id;
std::optional<json> message;
int64_t chat_id = 0;
bool private_chat = false;
std::optional<std::string> username;
};
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");
}
mw::Error invalidUpdate(std::string_view message)
{
return serviceError(502, "INVALID_UPDATE", message, std::nullopt, true);
}
bool getInt64(const json& value, int64_t& result)
{
if(!value.is_number_integer())
{
return false;
}
try
{
result = value.get<int64_t>();
return true;
}
catch(const std::exception&)
{
return false;
}
}
std::string lower(std::string value)
{
std::transform(value.begin(), value.end(), value.begin(),
[](unsigned char character)
{
if(character >= 'A' && character <= 'Z')
{
return static_cast<char>(character + ('a' - 'A'));
}
return static_cast<char>(character);
});
return value;
}
mw::E<std::vector<ValidatedUpdate>> validateUpdates(const json& updates)
{
if(!updates.is_array())
{
return std::unexpected(invalidUpdate("Telegram result is not an array"));
}
std::vector<ValidatedUpdate> result;
result.reserve(updates.size());
for(const json& update: updates)
{
if(!update.is_object() || !update.contains("update_id"))
{
return std::unexpected(invalidUpdate(
"Telegram update has no update_id"));
}
int64_t update_id = 0;
if(!getInt64(update["update_id"], update_id) || update_id < 0)
{
return std::unexpected(invalidUpdate(
"Telegram update_id is invalid"));
}
ValidatedUpdate item{update_id, std::nullopt, 0, false, std::nullopt};
if(update.contains("message"))
{
const json& message = update["message"];
if(!message.is_object() || !message.contains("chat") ||
!message["chat"].is_object() ||
!message["chat"].contains("id") ||
!getInt64(message["chat"]["id"], item.chat_id) ||
item.chat_id == 0)
{
return std::unexpected(invalidUpdate(
"Telegram message has an invalid chat"));
}
item.message = message;
if(message["chat"].contains("type") &&
message["chat"]["type"].is_string())
{
item.private_chat = message["chat"]["type"] == "private";
}
if(item.private_chat && message.contains("from") &&
message["from"].is_object() &&
message["from"].contains("username") &&
message["from"]["username"].is_string())
{
std::string value =
message["from"]["username"].get<std::string>();
if(!value.empty() && value.size() <= 64)
{
if(value.starts_with('@'))
{
value.erase(0, 1);
}
if(!value.empty())
{
item.username = lower(value);
}
}
}
}
result.push_back(std::move(item));
}
std::sort(result.begin(), result.end(),
[](const ValidatedUpdate& left, const ValidatedUpdate& right)
{
return left.update_id < right.update_id;
});
return result;
}
mw::E<std::optional<std::pair<int64_t, int64_t>>> pollState(
mw::SQLite& db)
{
auto rows = db.eval<int64_t, int64_t>(
"SELECT bot_id, next_offset FROM poll_state WHERE singleton = 1;");
if(!rows.has_value())
{
return std::unexpected(databaseError(rows.error()));
}
if(rows->empty())
{
return std::nullopt;
}
return std::pair{std::get<0>((*rows)[0]), std::get<1>((*rows)[0])};
}
mw::E<void> ensureBotInTransaction(mw::SQLite& db, int64_t bot_id)
{
auto state = pollState(db);
if(!state.has_value())
{
return std::unexpected(state.error());
}
if(state->has_value())
{
if(state->value().first != bot_id)
{
return std::unexpected(serviceError(
409, "BOT_MISMATCH",
"The state database belongs to another Telegram bot"));
}
return {};
}
auto statement = db.statementFromStr(
"INSERT INTO poll_state (singleton, bot_id, next_offset) "
"VALUES (1, ?, 0);");
if(!statement.has_value())
{
return std::unexpected(databaseError(statement.error()));
}
auto bind_result = statement->bind(bot_id);
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())
{
return std::unexpected(databaseError(insert_result.error()));
}
return {};
}
} // namespace
DeliveryStore::DeliveryStore(std::string db_path)
: db_path_(std::move(db_path))
{}
mw::E<void> DeliveryStore::ensureBot(int64_t bot_id) const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto begin_result = execute(**db_result, "BEGIN IMMEDIATE;");
if(!begin_result.has_value())
{
return std::unexpected(databaseError(begin_result.error()));
}
auto result = ensureBotInTransaction(**db_result, bot_id);
if(!result.has_value())
{
[[maybe_unused]] auto rollback = (*db_result)->execute("ROLLBACK;");
return std::unexpected(result.error());
}
auto commit_result = execute(**db_result, "COMMIT;");
if(!commit_result.has_value())
{
[[maybe_unused]] auto rollback = (*db_result)->execute("ROLLBACK;");
return std::unexpected(databaseError(commit_result.error()));
}
return {};
}
mw::E<int64_t> DeliveryStore::offset(int64_t bot_id) const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto state = pollState(**db_result);
if(!state.has_value())
{
return std::unexpected(state.error());
}
if(!state->has_value())
{
auto ensure_result = ensureBot(bot_id);
if(!ensure_result.has_value())
{
return std::unexpected(ensure_result.error());
}
return 0;
}
if(state->value().first != bot_id)
{
return std::unexpected(serviceError(
409, "BOT_MISMATCH",
"The state database belongs to another Telegram bot"));
}
return state->value().second;
}
mw::E<void> DeliveryStore::ingest(int64_t bot_id, const json& updates) const
{
auto validated = validateUpdates(updates);
if(!validated.has_value())
{
return std::unexpected(validated.error());
}
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 bot_result = ensureBotInTransaction(db, bot_id);
if(!bot_result.has_value())
{
rollback();
return std::unexpected(bot_result.error());
}
auto state = pollState(db);
if(!state.has_value())
{
rollback();
return std::unexpected(state.error());
}
int64_t next_offset = state->value().second;
for(const ValidatedUpdate& update: *validated)
{
if(update.update_id < next_offset)
{
continue;
}
if(update.update_id == std::numeric_limits<int64_t>::max())
{
rollback();
return std::unexpected(invalidUpdate("Telegram update_id overflow"));
}
if(update.message.has_value())
{
if(update.private_chat)
{
auto delete_username = db.statementFromStr(
"DELETE FROM usernames WHERE chat_id = ?;");
if(!delete_username.has_value())
{
rollback();
return std::unexpected(databaseError(
delete_username.error()));
}
auto bind_result = delete_username->bind(update.chat_id);
if(!bind_result.has_value())
{
rollback();
return std::unexpected(databaseError(bind_result.error()));
}
auto delete_result = db.execute(std::move(*delete_username));
if(!delete_result.has_value())
{
rollback();
return std::unexpected(databaseError(delete_result.error()));
}
if(update.username.has_value())
{
auto username_statement = db.statementFromStr(
"INSERT OR REPLACE INTO usernames "
"(username, chat_id, observed_at) VALUES (?, ?, ?);");
if(!username_statement.has_value())
{
rollback();
return std::unexpected(databaseError(
username_statement.error()));
}
auto username_bind = username_statement->bind(
*update.username, update.chat_id, nowSeconds());
if(!username_bind.has_value())
{
rollback();
return std::unexpected(databaseError(
username_bind.error()));
}
auto username_result =
db.execute(std::move(*username_statement));
if(!username_result.has_value())
{
rollback();
return std::unexpected(databaseError(
username_result.error()));
}
}
}
auto subscription_statement = db.statementFromStr(
"SELECT id, callback_url FROM subscriptions "
"WHERE chat_id = ?;");
if(!subscription_statement.has_value())
{
rollback();
return std::unexpected(databaseError(
subscription_statement.error()));
}
auto subscription_bind =
subscription_statement->bind(update.chat_id);
if(!subscription_bind.has_value())
{
rollback();
return std::unexpected(databaseError(subscription_bind.error()));
}
auto subscription_rows = db.eval<int64_t, std::string>(
std::move(*subscription_statement));
if(!subscription_rows.has_value())
{
rollback();
return std::unexpected(databaseError(subscription_rows.error()));
}
for(const auto& subscription: *subscription_rows)
{
auto delivery_statement = db.statementFromStr(
"INSERT OR IGNORE INTO deliveries "
"(subscription_id, update_id, payload, state, "
"attempt_count, next_attempt_at, created_at) "
"VALUES (?, ?, ?, 'PENDING', 0, ?, ?);");
if(!delivery_statement.has_value())
{
rollback();
return std::unexpected(databaseError(
delivery_statement.error()));
}
std::string payload = update.message->dump();
auto delivery_bind = delivery_statement->bind(
std::get<0>(subscription), update.update_id, payload,
nowSeconds(), nowSeconds());
if(!delivery_bind.has_value())
{
rollback();
return std::unexpected(databaseError(
delivery_bind.error()));
}
auto delivery_result = db.execute(std::move(*delivery_statement));
if(!delivery_result.has_value())
{
rollback();
return std::unexpected(databaseError(
delivery_result.error()));
}
auto queue_size = db.evalToValue<int64_t>(
"SELECT COUNT(*) FROM deliveries;");
if(!queue_size.has_value())
{
rollback();
return std::unexpected(databaseError(queue_size.error()));
}
if(*queue_size > 100000)
{
rollback();
return std::unexpected(serviceError(
503, "QUEUE_FULL", "The delivery queue is full",
std::nullopt, true));
}
}
}
next_offset = update.update_id + 1;
}
auto offset_statement = db.statementFromStr(
"UPDATE poll_state SET next_offset = ? WHERE singleton = 1;");
if(!offset_statement.has_value())
{
rollback();
return std::unexpected(databaseError(offset_statement.error()));
}
auto offset_bind = offset_statement->bind(next_offset);
if(!offset_bind.has_value())
{
rollback();
return std::unexpected(databaseError(offset_bind.error()));
}
auto offset_result = db.execute(std::move(*offset_statement));
if(!offset_result.has_value())
{
rollback();
return std::unexpected(databaseError(offset_result.error()));
}
auto commit_result = execute(db, "COMMIT;");
if(!commit_result.has_value())
{
rollback();
return std::unexpected(databaseError(commit_result.error()));
}
return {};
}
mw::E<std::optional<int64_t>> DeliveryStore::resolveUsername(
const std::string& username) const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto statement = (*db_result)->statementFromStr(
"SELECT chat_id FROM usernames WHERE username = ?;");
if(!statement.has_value())
{
return std::unexpected(databaseError(statement.error()));
}
auto bind_result = statement->bind(username);
if(!bind_result.has_value())
{
return std::unexpected(databaseError(bind_result.error()));
}
auto rows = (*db_result)->eval<int64_t>(std::move(*statement));
if(!rows.has_value())
{
return std::unexpected(databaseError(rows.error()));
}
if(rows->empty())
{
return std::nullopt;
}
return std::get<0>((*rows)[0]);
}
mw::E<std::optional<DeliveryJob>> DeliveryStore::claimNext() 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 rows = db.eval<int64_t, int64_t, int64_t, int64_t, int, std::string,
std::string, int64_t, int64_t, std::string>(
"SELECT d.id, p.bot_id, d.subscription_id, d.update_id, "
"d.attempt_count, s.callback_url, d.payload, "
"d.next_attempt_at, d.created_at, COALESCE(d.last_error, '') "
"FROM deliveries d "
"JOIN subscriptions s ON s.id = d.subscription_id "
"JOIN poll_state p ON p.singleton = 1 "
"WHERE d.state = 'PENDING' AND d.next_attempt_at <= "
"unixepoch() AND NOT EXISTS ("
"SELECT 1 FROM deliveries active "
"WHERE active.subscription_id = d.subscription_id "
"AND active.state = 'IN_FLIGHT') "
"ORDER BY d.next_attempt_at, d.id LIMIT 1;");
if(!rows.has_value())
{
rollback();
return std::unexpected(databaseError(rows.error()));
}
if(rows->empty())
{
auto commit_result = execute(db, "COMMIT;");
if(!commit_result.has_value())
{
rollback();
return std::unexpected(databaseError(commit_result.error()));
}
return std::nullopt;
}
const auto& row = (*rows)[0];
int64_t delivery_id = std::get<0>(row);
auto statement = db.statementFromStr(
"UPDATE deliveries SET state = 'IN_FLIGHT', "
"attempt_count = attempt_count + 1 "
"WHERE id = ? AND state = 'PENDING';");
if(!statement.has_value())
{
rollback();
return std::unexpected(databaseError(statement.error()));
}
auto bind_result = statement->bind(delivery_id);
if(!bind_result.has_value())
{
rollback();
return std::unexpected(databaseError(bind_result.error()));
}
auto update_result = db.execute(std::move(*statement));
if(!update_result.has_value())
{
rollback();
return std::unexpected(databaseError(update_result.error()));
}
auto commit_result = execute(db, "COMMIT;");
if(!commit_result.has_value())
{
rollback();
return std::unexpected(databaseError(commit_result.error()));
}
return DeliveryJob{delivery_id,
std::get<1>(row),
std::get<2>(row),
std::get<3>(row),
std::get<4>(row) + 1,
std::get<7>(row),
std::get<8>(row),
std::get<9>(row),
std::get<5>(row),
std::get<6>(row)};
}
mw::E<void> DeliveryStore::complete(int64_t delivery_id) const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto statement = (*db_result)->statementFromStr(
"DELETE FROM deliveries WHERE id = ? AND state = 'IN_FLIGHT';");
if(!statement.has_value())
{
return std::unexpected(databaseError(statement.error()));
}
auto bind_result = statement->bind(delivery_id);
if(!bind_result.has_value())
{
return std::unexpected(databaseError(bind_result.error()));
}
auto delete_result = (*db_result)->execute(std::move(*statement));
if(!delete_result.has_value())
{
return std::unexpected(databaseError(delete_result.error()));
}
return {};
}
mw::E<void> DeliveryStore::fail(int64_t delivery_id,
const std::string& error,
bool retryable,
std::optional<int> retry_after) 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 state_statement = db.statementFromStr(
"SELECT attempt_count, created_at FROM deliveries "
"WHERE id = ? AND state = 'IN_FLIGHT';");
if(!state_statement.has_value())
{
rollback();
return std::unexpected(databaseError(state_statement.error()));
}
auto bind_result = state_statement->bind(delivery_id);
if(!bind_result.has_value())
{
rollback();
return std::unexpected(databaseError(bind_result.error()));
}
auto rows = db.eval<int, int64_t>(std::move(*state_statement));
if(!rows.has_value())
{
rollback();
return std::unexpected(databaseError(rows.error()));
}
if(rows->empty())
{
auto commit_result = execute(db, "COMMIT;");
if(!commit_result.has_value())
{
rollback();
return std::unexpected(databaseError(commit_result.error()));
}
return {};
}
int attempts = std::get<0>((*rows)[0]);
int64_t created_at = std::get<1>((*rows)[0]);
int64_t current_time = nowSeconds();
constexpr int64_t MAX_DELIVERY_AGE = 24 * 60 * 60;
bool expired = current_time >= created_at &&
current_time - created_at >= MAX_DELIVERY_AGE;
bool dead = !retryable || attempts >= 8 || expired;
int64_t next_attempt = current_time;
if(!dead)
{
int exponent = std::min(attempts - 1, 8);
int64_t delay = 1LL << std::max(exponent, 0);
delay = std::min<int64_t>(delay, 300);
thread_local std::mt19937 random_generator(static_cast<unsigned>(
std::chrono::steady_clock::now().time_since_epoch().count()));
std::uniform_int_distribution<int64_t> jitter(
0, std::max<int64_t>(delay / 4, 1));
delay += jitter(random_generator);
if(retry_after.has_value())
{
delay = std::max<int64_t>(delay, *retry_after);
}
int64_t age = current_time >= created_at ?
current_time - created_at : 0;
int64_t remaining = MAX_DELIVERY_AGE -
std::min(age, MAX_DELIVERY_AGE);
if(remaining <= 0 || delay > remaining)
{
dead = true;
}
else
{
next_attempt += delay;
}
}
std::string safe_error = error.substr(0, 512);
auto update = db.statementFromStr(
"UPDATE deliveries SET state = ?, next_attempt_at = ?, "
"last_error = ? WHERE id = ? AND state = 'IN_FLIGHT';");
if(!update.has_value())
{
rollback();
return std::unexpected(databaseError(update.error()));
}
auto update_bind = update->bind(dead ? "DEAD" : "PENDING", next_attempt,
safe_error, delivery_id);
if(!update_bind.has_value())
{
rollback();
return std::unexpected(databaseError(update_bind.error()));
}
auto update_result = db.execute(std::move(*update));
if(!update_result.has_value())
{
rollback();
return std::unexpected(databaseError(update_result.error()));
}
auto commit_result = execute(db, "COMMIT;");
if(!commit_result.has_value())
{
rollback();
return std::unexpected(databaseError(commit_result.error()));
}
return {};
}
mw::E<void> DeliveryStore::resetInFlight() const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto result = (*db_result)->execute(
"UPDATE deliveries SET "
"state = CASE WHEN attempt_count >= 8 OR "
"created_at <= unixepoch() - 86400 THEN 'DEAD' ELSE 'PENDING' END, "
"next_attempt_at = unixepoch(), "
"last_error = CASE WHEN attempt_count >= 8 OR "
"created_at <= unixepoch() - 86400 "
"THEN COALESCE(last_error, 'Delivery budget exhausted after restart') "
"ELSE last_error END "
"WHERE state = 'IN_FLIGHT';");
if(!result.has_value())
{
return std::unexpected(databaseError(result.error()));
}
return {};
}
mw::E<int64_t> DeliveryStore::count() const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto result = (*db_result)->evalToValue<int64_t>(
"SELECT COUNT(*) FROM deliveries;");
if(!result.has_value())
{
return std::unexpected(databaseError(result.error()));
}
return *result;
}
mw::E<std::vector<DeliveryJob>> DeliveryStore::listDead() 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, int64_t, int64_t, int64_t, int,
int64_t, int64_t, std::string>(
"SELECT d.id, p.bot_id, d.subscription_id, d.update_id, "
"d.attempt_count, d.next_attempt_at, d.created_at, "
"COALESCE(d.last_error, '') "
"FROM deliveries d "
"JOIN poll_state p ON p.singleton = 1 WHERE d.state = 'DEAD' "
"ORDER BY d.id;");
if(!rows.has_value())
{
return std::unexpected(databaseError(rows.error()));
}
std::vector<DeliveryJob> 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), std::get<3>(row),
std::get<4>(row), std::get<5>(row),
std::get<6>(row), std::get<7>(row), {}, {}});
}
return result;
}
mw::E<bool> DeliveryStore::retry(int64_t delivery_id) const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto statement = (*db_result)->statementFromStr(
"UPDATE deliveries SET state = 'PENDING', attempt_count = 0, "
"next_attempt_at = unixepoch(), last_error = NULL "
"WHERE id = ? AND state = 'DEAD';");
if(!statement.has_value())
{
return std::unexpected(databaseError(statement.error()));
}
auto bind_result = statement->bind(delivery_id);
if(!bind_result.has_value())
{
return std::unexpected(databaseError(bind_result.error()));
}
auto result = (*db_result)->execute(std::move(*statement));
if(!result.has_value())
{
return std::unexpected(databaseError(result.error()));
}
return (*db_result)->changedRowsCount() > 0;
}
mw::E<bool> DeliveryStore::deleteDead(int64_t delivery_id) const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto statement = (*db_result)->statementFromStr(
"DELETE FROM deliveries WHERE id = ? AND state = 'DEAD';");
if(!statement.has_value())
{
return std::unexpected(databaseError(statement.error()));
}
auto bind_result = statement->bind(delivery_id);
if(!bind_result.has_value())
{
return std::unexpected(databaseError(bind_result.error()));
}
auto result = (*db_result)->execute(std::move(*statement));
if(!result.has_value())
{
return std::unexpected(databaseError(result.error()));
}
return (*db_result)->changedRowsCount() > 0;
}
mw::E<int64_t> DeliveryStore::purgeExpiredDead() const
{
auto db_result = openDatabase(db_path_);
if(!db_result.has_value())
{
return std::unexpected(db_result.error());
}
auto result = (*db_result)->execute(
"DELETE FROM deliveries WHERE state = 'DEAD' "
"AND next_attempt_at < unixepoch() - 604800;");
if(!result.has_value())
{
return std::unexpected(databaseError(result.error()));
}
return (*db_result)->changedRowsCount();
}
} // namespace telegrammer