use std::future::Future; use std::pin::Pin; use std::sync::atomic::{AtomicU64, Ordering}; use std::task::{Context, Poll}; use axum::{extract::Request, response::Response}; use tower::{Layer, Service}; use uuid::Uuid; /// Middleware that adds a unique X-Request-Id header to every request. #[derive(Clone, Default)] pub struct RequestIdLayer; impl Layer for RequestIdLayer { type Service = RequestIdMiddleware; fn layer(&self, inner: S) -> Self::Service { RequestIdMiddleware { inner, counter: AtomicU64::new(0), } } } pub struct RequestIdMiddleware { inner: S, counter: AtomicU64, } impl Clone for RequestIdMiddleware { fn clone(&self) -> Self { Self { inner: self.inner.clone(), counter: AtomicU64::new(self.counter.load(Ordering::Relaxed)), } } } impl Service> for RequestIdMiddleware where S: Service, Response = Response>, S::Future: Send + 'static, S::Error: 'static, ReqBody: Send + 'static, ResBody: Default + Send + 'static, { type Response = Response; type Error = S::Error; type Future = Pin> + Send>>; fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { self.inner.poll_ready(cx) } fn call(&mut self, req: Request) -> Self::Future { let request_id = Uuid::new_v4().to_string(); let (mut parts, body) = req.into_parts(); parts .headers .insert("x-request-id", request_id.parse().unwrap()); let req = Request::from_parts(parts, body); let fut = self.inner.call(req); Box::pin(async move { let mut response: Response = fut.await?; response .headers_mut() .insert("x-request-id", request_id.parse().unwrap()); Ok(response) }) } }