use crate::{ApiError, ApiErrorContext, ApiResult};
use anyhow::Result;
use axum::http::StatusCode;
pub trait ResultExt<T>: sealed::SealedResult {
fn context_status(
self,
status: StatusCode,
context: impl Into<ApiErrorContext>,
) -> ApiResult<T>;
fn context_bad_request(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_unauthorized(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_forbidden(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_not_found(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_method_not_allowed(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_conflict(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_unprocessable_entity(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_too_many_requests(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_internal(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_bad_gateway(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_service_unavailable(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_gateway_timeout(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
}
impl<T, E> ResultExt<T> for Result<T, E>
where
E: IntoApiError,
{
fn context_status(
self,
status: StatusCode,
context: impl Into<ApiErrorContext>,
) -> ApiResult<T> {
self.map_err(|err| err.context_status(status, context))
}
fn context_bad_request(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_bad_request(context))
}
fn context_unauthorized(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_unauthorized(context))
}
fn context_forbidden(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_forbidden(context))
}
fn context_not_found(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_not_found(context))
}
fn context_method_not_allowed(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_method_not_allowed(context))
}
fn context_conflict(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_conflict(context))
}
fn context_unprocessable_entity(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_unprocessable_entity(context))
}
fn context_too_many_requests(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_too_many_requests(context))
}
fn context_internal(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_internal(context))
}
fn context_bad_gateway(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_bad_gateway(context))
}
fn context_service_unavailable(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_service_unavailable(context))
}
fn context_gateway_timeout(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.map_err(|err| err.context_gateway_timeout(context))
}
}
pub trait OptionExt<T>: sealed::SealedOption {
fn context_status(
self,
status: StatusCode,
context: impl Into<ApiErrorContext>,
) -> ApiResult<T>;
fn context_bad_request(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_unauthorized(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_forbidden(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_not_found(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_method_not_allowed(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_conflict(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_unprocessable_entity(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_too_many_requests(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_internal(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_bad_gateway(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_service_unavailable(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
fn context_gateway_timeout(self, context: impl Into<ApiErrorContext>) -> ApiResult<T>;
}
impl<T> OptionExt<T> for Option<T> {
fn context_status(
self,
status: StatusCode,
context: impl Into<ApiErrorContext>,
) -> ApiResult<T> {
let ctx = context.into();
self.ok_or_else(|| {
let mut builder = ApiError::builder().status(status).title(ctx.title);
if let Some(detail) = ctx.detail {
builder = builder.detail(detail);
}
builder.build()
})
}
fn context_bad_request(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::BAD_REQUEST, context)
}
fn context_unauthorized(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::UNAUTHORIZED, context)
}
fn context_forbidden(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::FORBIDDEN, context)
}
fn context_not_found(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::NOT_FOUND, context)
}
fn context_method_not_allowed(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::METHOD_NOT_ALLOWED, context)
}
fn context_conflict(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::CONFLICT, context)
}
fn context_unprocessable_entity(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::UNPROCESSABLE_ENTITY, context)
}
fn context_too_many_requests(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::TOO_MANY_REQUESTS, context)
}
fn context_internal(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::INTERNAL_SERVER_ERROR, context)
}
fn context_bad_gateway(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::BAD_GATEWAY, context)
}
fn context_service_unavailable(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::SERVICE_UNAVAILABLE, context)
}
fn context_gateway_timeout(self, context: impl Into<ApiErrorContext>) -> ApiResult<T> {
self.context_status(StatusCode::GATEWAY_TIMEOUT, context)
}
}
pub trait IntoApiError: sealed::SealedIntoApiError {
fn context_status(self, status: StatusCode, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_bad_request(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_unauthorized(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_forbidden(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_not_found(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_method_not_allowed(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_conflict(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_unprocessable_entity(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_too_many_requests(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_internal(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_bad_gateway(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_service_unavailable(self, context: impl Into<ApiErrorContext>) -> ApiError;
fn context_gateway_timeout(self, context: impl Into<ApiErrorContext>) -> ApiError;
}
impl<E> IntoApiError for E
where
E: Into<anyhow::Error>,
{
fn context_status(self, status: StatusCode, context: impl Into<ApiErrorContext>) -> ApiError {
let ctx = context.into();
let mut builder = ApiError::builder()
.status(status)
.title(ctx.title)
.error(self);
if let Some(detail) = ctx.detail {
builder = builder.detail(detail);
}
builder.build()
}
fn context_bad_request(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::BAD_REQUEST, context)
}
fn context_unauthorized(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::UNAUTHORIZED, context)
}
fn context_forbidden(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::FORBIDDEN, context)
}
fn context_not_found(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::NOT_FOUND, context)
}
fn context_method_not_allowed(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::METHOD_NOT_ALLOWED, context)
}
fn context_conflict(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::CONFLICT, context)
}
fn context_unprocessable_entity(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::UNPROCESSABLE_ENTITY, context)
}
fn context_too_many_requests(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::TOO_MANY_REQUESTS, context)
}
fn context_internal(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::INTERNAL_SERVER_ERROR, context)
}
fn context_bad_gateway(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::BAD_GATEWAY, context)
}
fn context_service_unavailable(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::SERVICE_UNAVAILABLE, context)
}
fn context_gateway_timeout(self, context: impl Into<ApiErrorContext>) -> ApiError {
self.context_status(StatusCode::GATEWAY_TIMEOUT, context)
}
}
mod sealed {
use crate::IntoApiError;
pub trait SealedResult {}
pub trait SealedOption {}
pub trait SealedIntoApiError {}
impl<T, E> SealedResult for Result<T, E> where E: IntoApiError {}
impl<T> SealedOption for Option<T> {}
impl<E> SealedIntoApiError for E where E: Into<anyhow::Error> {}
}
#[cfg(test)]
mod tests {
use super::*;
use anyhow::anyhow;
#[test]
fn test_result_ext_context_bad_request_on_err() {
let result: Result<i32> = Err(anyhow!("Original error"));
let api_result = result.context_bad_request(("Bad Request", "Invalid data"));
assert!(api_result.is_err());
let err = api_result.unwrap_err();
assert_eq!(err.status(), StatusCode::BAD_REQUEST);
assert_eq!(err.title(), "Bad Request");
assert_eq!(err.detail(), Some("Invalid data"));
}
#[test]
fn test_result_ext_context_bad_request_title_only() {
let result: Result<i32> = Err(anyhow!("Original error"));
let api_result = result.context_bad_request("Bad Request");
assert!(api_result.is_err());
let err = api_result.unwrap_err();
assert_eq!(err.status(), StatusCode::BAD_REQUEST);
assert_eq!(err.title(), "Bad Request");
assert_eq!(err.detail(), None);
}
#[test]
fn test_result_ext_context_bad_request_on_ok() {
let result: Result<i32> = Ok(42);
let api_result = result.context_bad_request(("Bad Request", "Invalid data"));
assert!(api_result.is_ok());
assert_eq!(api_result.unwrap(), 42);
}
#[test]
fn test_result_ext_with_non_anyhow_error() {
let result = "not_a_number".parse::<i32>();
let api_result = result.context_bad_request(("Bad Request", "Value must be a number"));
assert!(api_result.is_err());
let err = api_result.unwrap_err();
assert_eq!(err.status(), StatusCode::BAD_REQUEST);
assert_eq!(err.title(), "Bad Request");
assert_eq!(err.detail(), Some("Value must be a number"));
}
#[test]
fn test_option_ext_context_bad_request_on_none() {
let option: Option<i32> = None;
let api_result = option.context_bad_request(("Bad Request", "Value is required"));
assert!(api_result.is_err());
let err = api_result.unwrap_err();
assert_eq!(err.status(), StatusCode::BAD_REQUEST);
assert_eq!(err.title(), "Bad Request");
assert_eq!(err.detail(), Some("Value is required"));
}
#[test]
fn test_option_ext_context_bad_request_title_only() {
let option: Option<i32> = None;
let api_result = option.context_bad_request("Bad Request");
assert!(api_result.is_err());
let err = api_result.unwrap_err();
assert_eq!(err.status(), StatusCode::BAD_REQUEST);
assert_eq!(err.title(), "Bad Request");
assert_eq!(err.detail(), None);
}
#[test]
fn test_option_ext_context_bad_request_on_some() {
let option: Option<i32> = Some(42);
let api_result = option.context_bad_request(("Bad Request", "Value is required"));
assert!(api_result.is_ok());
assert_eq!(api_result.unwrap(), 42);
}
#[test]
fn test_into_api_error_context_status() {
let anyhow_err = anyhow!("Custom error");
let api_err =
anyhow_err.context_status(StatusCode::IM_A_TEAPOT, ("Teapot", "I'm a teapot"));
assert_eq!(api_err.status(), StatusCode::IM_A_TEAPOT);
assert_eq!(api_err.title(), "Teapot");
assert_eq!(api_err.detail(), Some("I'm a teapot"));
}
#[test]
fn test_into_api_error_context_bad_request() {
let anyhow_err = anyhow!("Invalid input");
let api_err = anyhow_err.context_bad_request(("Bad Request", "Field validation failed"));
assert_eq!(api_err.status(), StatusCode::BAD_REQUEST);
assert_eq!(api_err.title(), "Bad Request");
assert_eq!(api_err.detail(), Some("Field validation failed"));
}
#[test]
fn test_into_api_error_title_only() {
let anyhow_err = anyhow!("Invalid input");
let api_err = anyhow_err.context_bad_request("Bad Request");
assert_eq!(api_err.status(), StatusCode::BAD_REQUEST);
assert_eq!(api_err.title(), "Bad Request");
assert_eq!(api_err.detail(), None);
}
#[test]
fn test_chaining_result_operations() {
fn get_value() -> Result<i32> {
Err(anyhow!("Failed to get value"))
}
let result = get_value().context_bad_request(("Bad Request", "Could not retrieve value"));
assert!(result.is_err());
assert_eq!(result.unwrap_err().status(), StatusCode::BAD_REQUEST);
}
#[test]
fn test_chaining_option_operations() {
fn get_value() -> Option<i32> {
None
}
let result = get_value().context_not_found(("Not Found", "Value does not exist"));
assert!(result.is_err());
assert_eq!(result.unwrap_err().status(), StatusCode::NOT_FOUND);
}
#[test]
fn test_question_mark_operator_with_result() {
fn helper() -> ApiResult<i32> {
let value: Result<i32> = Err(anyhow!("error"));
value.context_bad_request(("Bad Request", "Invalid"))?;
Ok(42)
}
let result = helper();
assert!(result.is_err());
assert_eq!(result.unwrap_err().status(), StatusCode::BAD_REQUEST);
}
#[test]
fn test_question_mark_operator_with_option() {
fn helper() -> ApiResult<i32> {
let value: Option<i32> = None;
value.context_not_found("Not Found")?;
Ok(42)
}
let result = helper();
assert!(result.is_err());
assert_eq!(result.unwrap_err().status(), StatusCode::NOT_FOUND);
}
#[test]
fn test_context_from_owned_strings() {
let result: Result<i32> = Err(anyhow!("error"));
let title = "Bad Request".to_string();
let detail = "Invalid input".to_string();
let api_result = result.context_bad_request((title, detail));
assert!(api_result.is_err());
let err = api_result.unwrap_err();
assert_eq!(err.title(), "Bad Request");
assert_eq!(err.detail(), Some("Invalid input"));
}
}