use axum::{
body::Body,
extract::Request,
http::header::{HeaderName, HeaderValue},
response::Response,
};
use std::{
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use tower::{Layer, Service};
use uuid::Uuid;
const X_REQUEST_ID: HeaderName = HeaderName::from_static("x-request-id");
#[derive(Clone)]
struct RequestId(String);
pub trait RequestIdGenerator {
const HEADER_NAME: HeaderName;
#[cfg(feature = "accept-client-id")]
const ID_LENGTH: usize;
fn generate(&self) -> HeaderValue;
}
#[derive(Clone, Debug)]
pub struct UuidGenerator;
impl RequestIdGenerator for UuidGenerator {
const HEADER_NAME: HeaderName = X_REQUEST_ID;
#[cfg(feature = "accept-client-id")]
const ID_LENGTH: usize = 36;
fn generate(&self) -> HeaderValue {
HeaderValue::from_str(&Uuid::new_v4().to_string())
.expect("UUIDv4 string is always valid ASCII")
}
}
impl Default for UuidGenerator {
fn default() -> Self {
Self
}
}
#[derive(Clone, Debug)]
pub struct RequestIdService<S, G> {
inner: S,
generator: Arc<G>,
}
impl<S, G> RequestIdService<S, G>
where
S: Clone,
G: RequestIdGenerator + Clone,
{
pub fn new(inner: S, generator: G) -> Self {
Self {
inner,
generator: generator.into(),
}
}
fn ensure_request_id(&self, req: &mut Request<Body>) -> HeaderValue {
#[cfg(feature = "accept-client-id")]
if let Some(existing_id) = req.headers().get(&G::HEADER_NAME) {
let existing_id = existing_id.clone();
if let Ok(id_str) = existing_id.to_str()
&& id_str.len() == G::ID_LENGTH
{
match req.extensions().get::<RequestId>() {
Some(ext) if ext.0 == id_str => return existing_id,
_ => {
req.extensions_mut().insert(RequestId(id_str.to_string()));
return existing_id;
}
}
}
}
let header_val = self.generator.generate();
req.headers_mut()
.insert(&G::HEADER_NAME, header_val.clone());
let request_id_str = header_val
.to_str()
.expect("RequestIdGenerator must produce ASCII header values")
.to_string();
req.extensions_mut().insert(RequestId(request_id_str));
header_val
}
}
impl<S, G> Service<Request<Body>> for RequestIdService<S, G>
where
S: Service<Request<Body>, Response = Response<Body>> + Send + Clone + 'static,
S::Future: Send + 'static,
S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
G: RequestIdGenerator + Send + Sync + Clone + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: Request<Body>) -> Self::Future {
let request_id = self.ensure_request_id(&mut req);
let fut = self.inner.call(req);
Box::pin(async move {
let mut res = fut.await?;
if res.headers().get(G::HEADER_NAME).is_none() {
res.headers_mut().insert(G::HEADER_NAME, request_id);
}
Ok(res)
})
}
}
#[derive(Clone, Debug)]
pub struct RequestIdLayer<G> {
generator: G,
}
impl<G> RequestIdLayer<G>
where
G: RequestIdGenerator,
{
pub fn new(generator: G) -> Self {
Self { generator }
}
}
impl Default for RequestIdLayer<UuidGenerator> {
fn default() -> Self {
RequestIdLayer {
generator: UuidGenerator,
}
}
}
impl<S, G> Layer<S> for RequestIdLayer<G>
where
G: RequestIdGenerator + Clone,
{
type Service = RequestIdService<S, G>;
fn layer(&self, service: S) -> Self::Service {
RequestIdService {
inner: service,
generator: Arc::new(self.generator.clone()),
}
}
}
pub trait RequestIdExt {
fn request_id(&self) -> Option<&str>;
}
impl RequestIdExt for Request<Body> {
fn request_id(&self) -> Option<&str> {
self.extensions().get::<RequestId>().map(|id| id.0.as_str())
}
}
#[cfg(test)]
mod tests;