use ::anyhow::Error as AnyhowError;
use ::axum::http::StatusCode;
use ::axum::response::IntoResponse;
use ::axum::response::Response;
use ::axum::Json;
use ::serde::Deserialize;
use ::serde::Serialize;
use ::std::fmt::Debug;
use ::std::fmt::Display;
use ::std::fmt::Formatter;
use ::std::fmt::Result as FmtResult;
use super::RouteErrorOutput;
pub struct RouteError<S = ()>
where
S: Serialize + for<'a> Deserialize<'a> + Debug,
{
status_code: StatusCode,
error: Option<AnyhowError>,
extra_data: Option<Box<S>>,
public_error_message: Option<String>,
}
impl RouteError<()> {
pub fn new_unauthorized() -> RouteError<()> {
Self::new_from_status(StatusCode::UNAUTHORIZED)
}
pub fn new_not_found() -> RouteError<()> {
Self::new_from_status(StatusCode::NOT_FOUND)
}
pub fn new_bad_request() -> RouteError<()> {
Self::new_from_status(StatusCode::BAD_REQUEST)
}
pub fn new_internal_server() -> RouteError<()> {
Self::new_from_status(StatusCode::INTERNAL_SERVER_ERROR)
}
pub fn new_conflict() -> RouteError<()> {
Self::new_from_status(StatusCode::CONFLICT)
}
pub fn new_from_status(status_code: StatusCode) -> RouteError<()> {
Self {
status_code,
..Self::default()
}
}
}
impl<S> RouteError<S>
where
S: Serialize + for<'a> Deserialize<'a> + Debug,
{
pub fn set_status_code(self, status_code: StatusCode) -> Self {
Self {
status_code,
..self
}
}
pub fn set_error(self, error: AnyhowError) -> Self {
Self {
error: Some(error),
..self
}
}
pub fn set_extra_data<NewS>(self, extra_data: NewS) -> RouteError<NewS>
where
NewS: Serialize + for<'a> Deserialize<'a> + Debug,
{
RouteError {
extra_data: Some(Box::new(extra_data)),
status_code: self.status_code,
error: self.error,
public_error_message: self.public_error_message,
}
}
pub fn set_public_error_message(self, public_error_message: &str) -> Self {
Self {
public_error_message: Some(public_error_message.to_string()),
..self
}
}
pub fn public_error_message<'a>(&'a self) -> &'a str {
if let Some(public_error_message) = self.public_error_message.as_ref() {
return public_error_message;
}
status_code_to_public_message(self.status_code())
}
pub fn status_code(&self) -> StatusCode {
self.status_code
}
}
impl<S> Default for RouteError<S>
where
S: Serialize + for<'a> Deserialize<'a> + Debug,
{
fn default() -> Self {
Self {
status_code: StatusCode::INTERNAL_SERVER_ERROR,
error: None,
extra_data: None,
public_error_message: None,
}
}
}
impl<S> IntoResponse for RouteError<S>
where
S: Serialize + for<'a> Deserialize<'a> + Debug,
{
fn into_response(self) -> Response {
let status = self.status_code();
let extra_data = self.extra_data;
let error = match self.public_error_message {
Some(public_error_message) => public_error_message,
None => status_code_to_public_message(status).to_string(),
};
let output = RouteErrorOutput { error, extra_data };
let body = Json(output);
(status, body).into_response()
}
}
impl Debug for RouteError {
fn fmt(&self, f: &mut Formatter<'_>) -> FmtResult {
write!(f, "{}, {:?}", self.public_error_message(), self.error)
}
}
impl Display for RouteError {
fn fmt(&self, f: &mut Formatter) -> FmtResult {
write!(f, "{}", self.public_error_message())
}
}
impl<S, FE> From<FE> for RouteError<S>
where
S: Serialize + for<'a> Deserialize<'a> + Debug,
FE: Into<AnyhowError>,
{
fn from(error: FE) -> Self {
RouteError {
status_code: StatusCode::INTERNAL_SERVER_ERROR,
error: Some(error.into()),
..Self::default()
}
}
}
fn status_code_to_public_message(status_code: StatusCode) -> &'static str {
match status_code {
StatusCode::CONFLICT => "The request is not allowed",
StatusCode::UNAUTHORIZED => "You are not authorised to access this endpoint",
StatusCode::NOT_FOUND => "The resource was not found",
StatusCode::BAD_REQUEST => "Bad request made",
_ => "An unexpected error occurred",
}
}