use axum::http::StatusCode;
use axum::response::{Html, IntoResponse, Response};
pub type Result<T = (), E = Error> = std::result::Result<T, E>;
#[non_exhaustive]
pub enum Error {
BadRequest(String),
Unauthorized,
Forbidden,
NotFound,
PageExpired,
TooManyRequests,
ServiceUnavailable,
Validation(crate::validation::ValidationError),
Status(StatusCode, String),
Internal(anyhow::Error),
}
pub fn abort(status: StatusCode, message: impl Into<String>) -> Error {
Error::Status(status, message.into())
}
pub fn abort_if(condition: bool, status: StatusCode, message: impl Into<String>) -> Result {
if condition {
Err(abort(status, message))
} else {
Ok(())
}
}
pub fn abort_unless(condition: bool, status: StatusCode, message: impl Into<String>) -> Result {
abort_if(!condition, status, message)
}
#[derive(Debug)]
struct Permanent(anyhow::Error);
impl std::fmt::Display for Permanent {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
std::fmt::Display::fmt(&self.0, f)
}
}
impl std::error::Error for Permanent {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
self.0.chain().nth(1)
}
}
impl Error {
pub fn permanent(err: impl Into<anyhow::Error>) -> Self {
Self::Internal(anyhow::Error::new(Permanent(err.into())))
}
pub fn permanent_message(message: impl std::fmt::Display) -> Self {
Self::permanent(anyhow::anyhow!("{message}"))
}
pub fn is_permanent(&self) -> bool {
match self {
Self::Internal(err) => err.chain().any(|e| e.is::<Permanent>()),
_ => false,
}
}
pub fn is_retryable(&self) -> bool {
match self {
Self::Internal(err) => err.chain().any(|e| {
e.downcast_ref::<crate::db::DbError>()
.is_some_and(crate::db::DbError::is_retryable)
}),
_ => false,
}
}
pub fn is_unique_violation(&self) -> bool {
match self {
Self::Internal(err) => err.chain().any(|e| {
e.downcast_ref::<crate::db::DbError>()
.is_some_and(crate::db::DbError::is_unique_violation)
|| matches!(
e.downcast_ref::<sqlx::Error>(),
Some(sqlx::Error::Database(db)) if db.is_unique_violation()
)
}),
_ => false,
}
}
pub fn status(&self) -> StatusCode {
match self {
Self::BadRequest(_) => StatusCode::BAD_REQUEST,
Self::Unauthorized => StatusCode::UNAUTHORIZED,
Self::Forbidden => StatusCode::FORBIDDEN,
Self::NotFound => StatusCode::NOT_FOUND,
Self::PageExpired => StatusCode::from_u16(419).expect("valid status code"),
Self::TooManyRequests => StatusCode::TOO_MANY_REQUESTS,
Self::ServiceUnavailable => StatusCode::SERVICE_UNAVAILABLE,
Self::Validation(_) => StatusCode::UNPROCESSABLE_ENTITY,
Self::Status(status, _) => *status,
Self::Internal(_) if self.is_unique_violation() => StatusCode::CONFLICT,
Self::Internal(_) => StatusCode::INTERNAL_SERVER_ERROR,
}
}
}
impl std::fmt::Debug for Error {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Internal(err) => write!(f, "{err:?}"),
Self::BadRequest(msg) => write!(f, "bad request: {msg}"),
Self::Status(status, msg) => write!(f, "{}: {msg}", status.as_u16()),
Self::Validation(err) => write!(f, "validation failed: {:?}", err.errors),
other => write!(f, "{}", reason(other.status())),
}
}
}
impl<E: Into<anyhow::Error>> From<E> for Error {
fn from(err: E) -> Self {
Self::Internal(err.into())
}
}
#[derive(Debug, Clone)]
pub(crate) struct ErrorPage {
pub status: StatusCode,
pub detail: Option<String>,
pub debug_detail: Option<String>,
pub template: Option<String>,
}
impl ErrorPage {
pub fn shown_detail(&self, debug: bool) -> Option<&str> {
self.detail.as_deref().or(if debug {
self.debug_detail.as_deref()
} else {
None
})
}
}
impl ErrorPage {
pub fn json(&self, debug: bool) -> Response {
let message = self
.shown_detail(debug)
.unwrap_or_else(|| reason(self.status));
(
self.status,
axum::Json(serde_json::json!({ "message": message })),
)
.into_response()
}
}
pub(crate) fn wants_json(headers: &axum::http::HeaderMap) -> bool {
let header = |name| headers.get(name).and_then(|v| v.to_str().ok());
header(axum::http::header::ACCEPT).is_some_and(|v| v.contains("application/json"))
|| header(axum::http::header::CONTENT_TYPE)
.is_some_and(|v| v.starts_with("application/json"))
}
pub(crate) fn panic_message(panic: &(dyn std::any::Any + Send)) -> String {
panic
.downcast_ref::<String>()
.cloned()
.or_else(|| panic.downcast_ref::<&str>().map(|s| (*s).to_owned()))
.unwrap_or_else(|| "(no message)".into())
}
pub(crate) fn reason(status: StatusCode) -> &'static str {
match status.as_u16() {
419 => "Page Expired",
_ => status.canonical_reason().unwrap_or("Error"),
}
}
impl From<crate::validation::ValidationError> for Error {
fn from(err: crate::validation::ValidationError) -> Self {
Self::Validation(err)
}
}
impl From<crate::validation::Errors> for Error {
fn from(errors: crate::validation::Errors) -> Self {
Self::Validation(errors.into())
}
}
impl IntoResponse for Error {
fn into_response(self) -> Response {
if let Self::Validation(err) = self {
return err.into_response();
}
let status = self.status();
let mut template = None;
let (detail, debug_detail) = match &self {
Self::Internal(err) => {
tracing::error!(error = ?err, "internal server error");
crate::report::request_error(err);
template = err
.chain()
.find_map(|e| e.downcast_ref::<minijinja::Error>())
.map(|e| e.display_debug_info().to_string())
.filter(|info| !info.trim().is_empty());
(None, Some(format!("{err:?}")))
}
Self::BadRequest(msg) => (Some(msg.clone()), None),
Self::Status(_, msg) if !msg.is_empty() => (Some(msg.clone()), None),
_ => (None, None),
};
let mut res = (status, Html(error_page(status, detail.as_deref()))).into_response();
res.extensions_mut().insert(ErrorPage {
status,
detail,
debug_detail,
template,
});
res
}
}
pub(crate) fn error_page(status: StatusCode, detail: Option<&str>) -> String {
let code = status.as_u16();
let reason = reason(status);
let detail = detail
.map(|d| format!("<pre>{}</pre>", escape(d)))
.unwrap_or_default();
format!(
"<!doctype html><html><head><meta charset=\"utf-8\">\
<meta name=\"viewport\" content=\"width=device-width, initial-scale=1\">\
<title>{code} {reason}</title></head>\
<body><h1>{code} · {reason}</h1>{detail}</body></html>"
)
}
fn escape(s: &str) -> String {
s.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validation_errors_answer_422_and_other_errors_500() {
let mut errors = crate::validation::Errors::new();
errors.add("name", "The name field is required.");
let err = Error::Validation(crate::validation::ValidationError::new(errors));
assert_eq!(err.status(), StatusCode::UNPROCESSABLE_ENTITY);
let plain = Error::from(anyhow::anyhow!("no database involved"));
assert!(!plain.is_unique_violation());
assert_eq!(plain.status(), StatusCode::INTERNAL_SERVER_ERROR);
}
#[tokio::test]
async fn a_raw_sqlx_unique_violation_is_a_conflict() {
let db = crate::db::connect(&crate::Config::default()).await.unwrap();
let Some(pool) = db.sqlite() else {
return; };
sqlx::raw_sql("CREATE TABLE tags (name TEXT UNIQUE); INSERT INTO tags VALUES ('a')")
.execute(pool)
.await
.unwrap();
let raw = sqlx::raw_sql("INSERT INTO tags VALUES ('a')")
.execute(pool)
.await
.unwrap_err();
let err = Error::from(anyhow::Error::new(raw));
assert!(err.is_unique_violation());
assert_eq!(err.status(), StatusCode::CONFLICT);
}
}