use crate::ApiErrorBuilder;
use axum::{
extract::Request,
http::{HeaderMap, Method, Uri},
response::Response,
};
use futures_util::future::BoxFuture;
use std::{
cell::RefCell,
sync::Arc,
task::{Context, Poll},
};
use tower::{Layer, Service};
thread_local! {
static ENRICHMENT_CONTEXT: RefCell<Option<EnrichmentContext>> = const { RefCell::new(None) };
}
#[derive(Clone, Debug)]
pub struct RequestSnapshot {
method: Method,
uri: Uri,
headers: HeaderMap,
}
impl RequestSnapshot {
pub fn method(&self) -> &Method {
&self.method
}
pub fn uri(&self) -> &Uri {
&self.uri
}
pub fn headers(&self) -> &HeaderMap {
&self.headers
}
pub fn from_request(request: &Request) -> Self {
Self {
method: request.method().clone(),
uri: request.uri().clone(),
headers: request.headers().clone(),
}
}
}
type ErrorEnricher =
Arc<dyn Fn(ApiErrorBuilder, &RequestSnapshot) -> ApiErrorBuilder + Send + Sync + 'static>;
#[derive(Clone)]
pub(crate) struct EnrichmentContext {
request: RequestSnapshot,
enricher: ErrorEnricher,
}
impl EnrichmentContext {
fn new(request: RequestSnapshot, enricher: ErrorEnricher) -> Self {
Self { request, enricher }
}
fn set(self) {
ENRICHMENT_CONTEXT.with(|data| {
*data.borrow_mut() = Some(self);
});
}
fn clear() {
ENRICHMENT_CONTEXT.with(|data| {
*data.borrow_mut() = None;
});
}
fn apply(&self, builder: ApiErrorBuilder) -> ApiErrorBuilder {
(self.enricher)(builder, &self.request)
}
pub(crate) fn invoke(builder: ApiErrorBuilder) -> ApiErrorBuilder {
ENRICHMENT_CONTEXT.with(|data| {
if let Some(enrichment_ctx) = data.borrow().as_ref() {
enrichment_ctx.apply(builder)
} else {
builder
}
})
}
}
pub struct ErrorInterceptor<S> {
inner: S,
enricher: ErrorEnricher,
}
impl<S> Clone for ErrorInterceptor<S>
where
S: Clone,
{
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
enricher: self.enricher.clone(),
}
}
}
impl<S> Service<Request> for ErrorInterceptor<S>
where
S: Service<Request, Response = Response> + Send + 'static,
S::Future: Send + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, request: Request) -> Self::Future {
let snapshot = RequestSnapshot::from_request(&request);
let ctx = EnrichmentContext::new(snapshot, self.enricher.clone());
let future = self.inner.call(request);
Box::pin(async move {
ctx.set();
let result = future.await;
EnrichmentContext::clear();
result
})
}
}
#[derive(Clone)]
pub struct ErrorInterceptorLayer {
enricher: ErrorEnricher,
}
impl ErrorInterceptorLayer {
pub fn new<F>(enricher: F) -> Self
where
F: Fn(ApiErrorBuilder, &RequestSnapshot) -> ApiErrorBuilder + Send + Sync + 'static,
{
Self {
enricher: Arc::new(enricher),
}
}
}
impl<S> Layer<S> for ErrorInterceptorLayer {
type Service = ErrorInterceptor<S>;
fn layer(&self, inner: S) -> Self::Service {
ErrorInterceptor {
inner,
enricher: self.enricher.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
use serde_json::json;
use serial_test::serial;
#[test]
#[serial]
fn test_error_enricher() {
let enricher = Arc::new(|builder: ApiErrorBuilder, req: &RequestSnapshot| {
builder.meta(json!({
"method": req.method.as_str(),
"uri": req.uri.to_string(),
}))
});
let snapshot = RequestSnapshot {
method: Method::GET,
uri: "/test".parse().unwrap(),
headers: HeaderMap::default(),
};
EnrichmentContext::new(snapshot, enricher).set();
let error = crate::ApiError::builder()
.status(StatusCode::NOT_FOUND)
.title("Not Found")
.detail("Resource not found")
.build();
assert!(error.meta().is_some());
let meta = error.meta().unwrap();
assert_eq!(meta["method"], "GET");
assert_eq!(meta["uri"], "/test");
EnrichmentContext::clear();
}
#[test]
#[serial]
fn test_enricher_without_context() {
EnrichmentContext::clear();
let error = crate::ApiError::builder()
.status(StatusCode::BAD_REQUEST)
.title("Bad Request")
.detail("Invalid input")
.build();
assert!(error.meta().is_none());
}
#[test]
#[serial]
fn test_request_data_lifecycle() {
let snapshot = RequestSnapshot {
method: Method::POST,
uri: "/api/users".parse().unwrap(),
headers: HeaderMap::default(),
};
let enricher = Arc::new(|builder: ApiErrorBuilder, _req: &RequestSnapshot| builder);
EnrichmentContext::new(snapshot.clone(), enricher).set();
ENRICHMENT_CONTEXT.with(|data| {
let borrowed = data.borrow();
assert!(borrowed.is_some());
let stored_req = &borrowed.as_ref().unwrap().request;
assert_eq!(stored_req.method, Method::POST);
assert_eq!(stored_req.uri.to_string(), "/api/users");
});
EnrichmentContext::clear();
ENRICHMENT_CONTEXT.with(|data| {
assert!(data.borrow().is_none());
});
}
}