macro_rules! try_api {
($expr:expr, $context:expr) => {
match $expr {
Ok(value) => value,
Err(e) => return Ok($crate::api::internal_error($context, &e)),
}
};
}
pub mod archive;
pub mod history;
pub mod jobs;
pub mod queues;
pub mod stats;
pub mod system;
use serde::{Deserialize, Serialize};
use warp::Reply;
#[derive(Debug, Serialize)]
pub struct ApiResponse<T> {
pub success: bool,
pub data: Option<T>,
pub error: Option<String>,
pub timestamp: chrono::DateTime<chrono::Utc>,
}
impl<T> ApiResponse<T> {
pub fn success(data: T) -> Self {
Self {
success: true,
data: Some(data),
error: None,
timestamp: chrono::Utc::now(),
}
}
pub fn error(message: String) -> Self {
Self {
success: false,
data: None,
error: Some(message),
timestamp: chrono::Utc::now(),
}
}
}
pub fn json_reply<T: Serialize>(body: &T) -> warp::reply::Response {
warp::reply::with_status(warp::reply::json(body), warp::http::StatusCode::OK).into_response()
}
pub fn error_reply(
status: warp::http::StatusCode,
message: impl Into<String>,
) -> warp::reply::Response {
let body = ApiResponse::<()>::error(message.into());
warp::reply::with_status(warp::reply::json(&body), status).into_response()
}
pub fn internal_error(context: &str, err: &dyn std::fmt::Display) -> warp::reply::Response {
tracing::error!(error = %err, "{}", context);
error_reply(
warp::http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{}: {}", context, err),
)
}
pub fn queue_error_reply(
context: &str,
err: &hammerwork::HammerworkError,
) -> warp::reply::Response {
use hammerwork::HammerworkError;
match err {
HammerworkError::JobNotFound { .. } => error_reply(
warp::http::StatusCode::NOT_FOUND,
format!("{}: {}", context, err),
),
HammerworkError::InvalidJobTransition { .. } => error_reply(
warp::http::StatusCode::CONFLICT,
format!("{}: {}", context, err),
),
_ => internal_error(context, err),
}
}
#[derive(Debug, Deserialize)]
pub struct PaginationParams {
pub page: Option<u32>,
pub limit: Option<u32>,
pub offset: Option<u32>,
}
impl Default for PaginationParams {
fn default() -> Self {
Self {
page: Some(1),
limit: Some(50),
offset: None,
}
}
}
pub const DEFAULT_PAGE_SIZE: u32 = 50;
pub const MAX_PAGE_SIZE: u32 = 1000;
impl PaginationParams {
pub fn get_offset(&self) -> u32 {
if let Some(offset) = self.offset {
offset
} else {
let page = self.page.unwrap_or(1);
page.saturating_sub(1).saturating_mul(self.get_limit())
}
}
pub fn get_limit(&self) -> u32 {
self.limit
.unwrap_or(DEFAULT_PAGE_SIZE)
.clamp(1, MAX_PAGE_SIZE)
}
}
#[derive(Debug, Serialize)]
pub struct PaginationMeta {
pub page: u32,
pub limit: u32,
pub offset: u32,
pub total: u64,
pub total_pages: u32,
pub has_next: bool,
pub has_prev: bool,
}
impl PaginationMeta {
pub fn new(params: &PaginationParams, total: u64) -> Self {
let limit = params.get_limit();
let offset = params.get_offset();
let page = params.page.unwrap_or(1);
let total_pages = u32::try_from(total.div_ceil(u64::from(limit))).unwrap_or(u32::MAX);
Self {
page,
limit,
offset,
total,
total_pages,
has_next: page < total_pages,
has_prev: page > 1,
}
}
}
#[derive(Debug, Serialize)]
pub struct PaginatedResponse<T> {
pub items: Vec<T>,
pub pagination: PaginationMeta,
}
#[derive(Debug, Deserialize, Default)]
pub struct FilterParams {
pub status: Option<String>,
pub priority: Option<String>,
pub queue: Option<String>,
pub created_after: Option<chrono::DateTime<chrono::Utc>>,
pub created_before: Option<chrono::DateTime<chrono::Utc>>,
pub search: Option<String>,
}
#[derive(Debug, Deserialize, Default)]
pub struct SortParams {
pub sort_by: Option<String>,
pub sort_order: Option<String>,
}
impl SortParams {
pub fn get_order_by(&self) -> (String, String) {
let field = self.sort_by.as_deref().unwrap_or("created_at").to_string();
let direction = match self.sort_order.as_deref() {
Some("asc") | Some("ASC") => "ASC".to_string(),
_ => "DESC".to_string(),
};
(field, direction)
}
}
pub async fn handle_api_error(
err: warp::Rejection,
) -> Result<impl warp::Reply, std::convert::Infallible> {
let response = if err.is_not_found() {
ApiResponse::<()>::error("Resource not found".to_string())
} else if err
.find::<warp::filters::body::BodyDeserializeError>()
.is_some()
{
ApiResponse::<()>::error("Invalid request body".to_string())
} else if err.find::<warp::reject::InvalidQuery>().is_some() {
ApiResponse::<()>::error("Invalid query parameters".to_string())
} else {
ApiResponse::<()>::error("Internal server error".to_string())
};
let status = if err.is_not_found() {
warp::http::StatusCode::NOT_FOUND
} else if err
.find::<warp::filters::body::BodyDeserializeError>()
.is_some()
|| err.find::<warp::reject::InvalidQuery>().is_some()
{
warp::http::StatusCode::BAD_REQUEST
} else {
warp::http::StatusCode::INTERNAL_SERVER_ERROR
};
Ok(warp::reply::with_status(
warp::reply::json(&response),
status,
))
}
pub fn with_pagination()
-> impl warp::Filter<Extract = (PaginationParams,), Error = warp::Rejection> + Clone {
warp::query::<PaginationParams>()
}
pub fn with_filters()
-> impl warp::Filter<Extract = (FilterParams,), Error = warp::Rejection> + Clone {
warp::query::<FilterParams>()
}
pub fn with_sort() -> impl warp::Filter<Extract = (SortParams,), Error = warp::Rejection> + Clone {
warp::query::<SortParams>()
}
#[cfg(test)]
pub(crate) mod test_support {
use hammerwork::JobQueue;
use std::sync::Arc;
pub fn unreachable_queue() -> Arc<JobQueue<sqlx::Postgres>> {
let pool = sqlx::postgres::PgPoolOptions::new()
.acquire_timeout(std::time::Duration::from_millis(200))
.connect_lazy("postgres://nobody:nothing@127.0.0.1:1/none")
.expect("lazy pool");
Arc::new(JobQueue::new(pool))
}
pub async fn body_json(response: warp::reply::Response) -> (u16, serde_json::Value) {
use warp::Filter;
let slot = Arc::new(std::sync::Mutex::new(Some(response)));
let filter = warp::any().map(move || slot.lock().unwrap().take().unwrap());
let reply = warp::test::request().reply(&filter).await;
(
reply.status().as_u16(),
serde_json::from_slice(reply.body()).unwrap(),
)
}
}
#[cfg(test)]
mod tests {
use super::test_support::body_json;
use super::*;
#[test]
fn test_api_response_success() {
let response = ApiResponse::success("test data");
assert!(response.success);
assert_eq!(response.data, Some("test data"));
assert!(response.error.is_none());
}
#[test]
fn test_api_response_error() {
let response: ApiResponse<()> = ApiResponse::error("Something went wrong".to_string());
assert!(!response.success);
assert!(response.data.is_none());
assert_eq!(response.error, Some("Something went wrong".to_string()));
}
#[test]
fn test_pagination_params_defaults() {
let params = PaginationParams::default();
assert_eq!(params.get_limit(), 50);
assert_eq!(params.get_offset(), 0);
}
#[test]
fn test_pagination_params_calculation() {
let params = PaginationParams {
page: Some(3),
limit: Some(20),
offset: None,
};
assert_eq!(params.get_limit(), 20);
assert_eq!(params.get_offset(), 40); }
#[test]
fn the_offset_uses_the_clamped_limit_and_never_overflows() {
let params = PaginationParams {
page: Some(2),
limit: Some(5000),
offset: None,
};
assert_eq!(params.get_limit(), MAX_PAGE_SIZE);
assert_eq!(params.get_offset(), MAX_PAGE_SIZE);
let huge = PaginationParams {
page: Some(u32::MAX),
limit: Some(u32::MAX),
offset: None,
};
assert_eq!(
huge.get_offset(),
u32::MAX,
"saturates instead of overflowing"
);
let zero = PaginationParams {
page: Some(0),
limit: Some(0),
offset: None,
};
assert_eq!(zero.get_limit(), 1, "a zero limit is clamped to one");
assert_eq!(zero.get_offset(), 0);
let meta = PaginationMeta::new(&zero, 3);
assert_eq!(meta.total_pages, 3);
let explicit = PaginationParams {
page: Some(7),
limit: Some(10),
offset: Some(5),
};
assert_eq!(explicit.get_offset(), 5, "an explicit offset wins");
}
#[test]
fn test_pagination_meta() {
let params = PaginationParams {
page: Some(2),
limit: Some(10),
offset: None,
};
let meta = PaginationMeta::new(¶ms, 45);
assert_eq!(meta.page, 2);
assert_eq!(meta.limit, 10);
assert_eq!(meta.total, 45);
assert_eq!(meta.total_pages, 5);
assert!(meta.has_next);
assert!(meta.has_prev);
}
#[test]
fn test_sort_params_defaults() {
let params = SortParams {
sort_by: None,
sort_order: None,
};
let (field, direction) = params.get_order_by();
assert_eq!(field, "created_at");
assert_eq!(direction, "DESC");
}
#[test]
fn test_sort_params_custom() {
let params = SortParams {
sort_by: Some("name".to_string()),
sort_order: Some("asc".to_string()),
};
let (field, direction) = params.get_order_by();
assert_eq!(field, "name");
assert_eq!(direction, "ASC");
}
#[tokio::test]
async fn test_internal_error_is_500_with_json_error_body() {
let (status, body) = body_json(internal_error("Failed to list jobs", &"db down")).await;
assert_eq!(status, 500);
assert_eq!(body["success"], false);
assert_eq!(body["error"], "Failed to list jobs: db down");
assert!(body["data"].is_null());
}
#[tokio::test]
async fn test_error_reply_uses_given_status() {
let (status, body) =
body_json(error_reply(warp::http::StatusCode::BAD_REQUEST, "nope")).await;
assert_eq!(status, 400);
assert_eq!(body["error"], "nope");
}
#[tokio::test]
async fn test_json_reply_is_200() {
let (status, body) = body_json(json_reply(&ApiResponse::success(5))).await;
assert_eq!(status, 200);
assert_eq!(body["data"], 5);
}
}
#[cfg(test)]
mod undecodable_row_tests {
use super::history::JobHistory;
use hammerwork::archive::{ArchivalConfig, ArchivalPolicy, ArchivalReason};
use hammerwork::{Job, JobId};
use serde_json::{Value, json};
use std::future::Future;
use std::sync::Arc;
use warp::Filter;
fn listed(body: &Value) -> Vec<String> {
let mut ids: Vec<String> = body["data"]["items"]
.as_array()
.unwrap_or_else(|| panic!("no items: {body}"))
.iter()
.map(|item| item["id"].as_str().unwrap().to_string())
.collect();
ids.sort();
ids
}
async fn listings_skip_undecodable_rows<Q, S, F>(
queue: Arc<Q>,
sql: S,
corrupt_archive_id: bool,
) where
Q: JobHistory + 'static,
S: Fn(String) -> F,
F: Future<Output = ()>,
{
let routes = super::jobs::routes(queue.clone())
.or(super::queues::routes(queue.clone()))
.or(super::archive::archive_routes(queue.clone()));
let get = |path: String| {
let routes = routes.clone();
async move {
let response = warp::test::request().path(&path).reply(&routes).await;
let body: Value = serde_json::from_slice(response.body()).unwrap();
(response.status().as_u16(), body)
}
};
let tag = uuid::Uuid::new_v4().simple().to_string();
let name = format!("undecodable_{tag}");
let archive_name = format!("undecodable_archive_{tag}");
let enqueue = |queue_name: &str| queue.enqueue(Job::new(queue_name.to_string(), json!({})));
let break_timeout = |id: JobId| {
sql(format!(
"UPDATE hammerwork_jobs SET timeout_seconds = -1 WHERE id = '{id}'"
))
};
let ready = enqueue(&name).await.unwrap();
let bad_ready = enqueue(&name).await.unwrap();
break_timeout(bad_ready).await;
let dead = enqueue(&name).await.unwrap();
let bad_dead = enqueue(&name).await.unwrap();
for id in [dead, bad_dead] {
queue.mark_job_dead(id, "gone").await.unwrap();
}
break_timeout(bad_dead).await;
let mut healthy = vec![ready.to_string(), dead.to_string()];
healthy.sort();
for path in [
format!("/jobs?queue={name}&limit=100"),
format!("/queues/{name}/jobs?limit=100"),
] {
let (status, body) = get(path.clone()).await;
assert_eq!(status, 200, "{path}: {body}");
assert_eq!(listed(&body), healthy, "{path}");
}
let response = warp::test::request()
.method("POST")
.path("/jobs/search?limit=100")
.json(&json!({"query": "", "queues": [name]}))
.reply(&routes)
.await;
let body: Value = serde_json::from_slice(response.body()).unwrap();
assert_eq!(response.status(), 200, "{body}");
assert_eq!(listed(&body), healthy, "search");
let (status, body) = get(format!("/jobs/{bad_ready}")).await;
assert_eq!(status, 500, "{body}");
assert!(
body["error"].as_str().unwrap().contains("timeout_seconds"),
"{body}"
);
let archived = enqueue(&archive_name).await.unwrap();
let bad_archived = enqueue(&archive_name).await.unwrap();
for id in [archived, bad_archived] {
queue.mark_job_dead(id, "gone").await.unwrap();
}
let policy = ArchivalPolicy::new()
.archive_dead_after(chrono::Duration::seconds(0))
.enabled(true);
queue
.archive_jobs(
Some(&archive_name),
&policy,
&ArchivalConfig::new(),
ArchivalReason::Manual,
None,
)
.await
.unwrap();
if corrupt_archive_id {
sql(format!(
"UPDATE hammerwork_jobs_archive SET id = 'not-a-uuid' WHERE id = '{bad_archived}'"
))
.await;
}
let (status, body) = get(format!("/archive/jobs?queue={archive_name}")).await;
assert_eq!(status, 200, "{body}");
let ids = listed(&body);
assert!(ids.contains(&archived.to_string()), "{body}");
if corrupt_archive_id {
assert_eq!(ids, [archived.to_string()], "{body}");
}
for (table, queue_name) in [
("hammerwork_jobs", &name),
("hammerwork_jobs_archive", &archive_name),
] {
sql(format!(
"DELETE FROM {table} WHERE queue_name = '{queue_name}'"
))
.await;
}
}
#[tokio::test]
#[ignore = "requires DATABASE_URL (PostgreSQL)"]
async fn postgres_listings_skip_undecodable_rows() {
let url = std::env::var("DATABASE_URL").expect("DATABASE_URL");
let pool = sqlx::PgPool::connect(&url).await.unwrap();
let queue = Arc::new(hammerwork::JobQueue::new(pool.clone()));
let sql = |statement: String| {
let pool = pool.clone();
async move {
sqlx::query(&statement).execute(&pool).await.unwrap();
}
};
listings_skip_undecodable_rows(queue, sql, false).await;
}
#[tokio::test]
#[ignore = "requires MYSQL_DATABASE_URL"]
async fn mysql_listings_skip_undecodable_rows() {
let url = std::env::var("MYSQL_DATABASE_URL").expect("MYSQL_DATABASE_URL");
let pool = sqlx::MySqlPool::connect(&url).await.unwrap();
let queue = Arc::new(hammerwork::JobQueue::new(pool.clone()));
let sql = |statement: String| {
let pool = pool.clone();
async move {
sqlx::query(&statement).execute(&pool).await.unwrap();
}
};
listings_skip_undecodable_rows(queue, sql, true).await;
}
}