use tower::{Layer, Service};
use uuid::Uuid;
pub const REQUEST_ID_HEADER: &str = "x-request-id";
const MAX_REQUEST_ID_LEN: usize = 128;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RequestId(pub String);
impl RequestId {
pub fn extract<B>(req: &http::Request<B>) -> Option<String> {
req.extensions()
.get::<RequestId>()
.map(|id| id.0.clone())
.or_else(|| get_request_id(req.headers()))
}
}
impl std::fmt::Display for RequestId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct RequestIdLayer;
impl<S> Layer<S> for RequestIdLayer {
type Service = RequestIdService<S>;
fn layer(&self, inner: S) -> Self::Service {
RequestIdService { inner }
}
}
#[derive(Debug, Clone)]
pub struct RequestIdService<S> {
inner: S,
}
fn incoming_request_id(headers: &http::HeaderMap) -> Option<String> {
let value = headers.get(REQUEST_ID_HEADER)?.to_str().ok()?;
if value.is_empty() || value.len() > MAX_REQUEST_ID_LEN {
return None;
}
Some(value.to_string())
}
fn generated_request_id() -> (String, Option<http::HeaderValue>) {
let id = Uuid::new_v4().to_string();
let value = http::HeaderValue::from_str(&id).ok();
(id, value)
}
impl<S, B, ResB> Service<http::Request<B>> for RequestIdService<S>
where
S: Service<http::Request<B>, Response = http::Response<ResB>> + Send + 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
B: Send + 'static,
ResB: Send + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>,
>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: http::Request<B>) -> Self::Future {
let (request_id, header_value) = match incoming_request_id(req.headers()) {
Some(id) => match http::HeaderValue::from_str(&id) {
Ok(value) => (id, Some(value)),
Err(e) => {
#[cfg(feature = "tracing")]
tracing::debug!(msg = "replacing unusable client request id", error = %e);
let _ = e;
generated_request_id()
}
},
None => generated_request_id(),
};
let Some(header_value) = header_value else {
return Box::pin(self.inner.call(req));
};
req.headers_mut()
.insert(REQUEST_ID_HEADER, header_value.clone());
req.extensions_mut().insert(RequestId(request_id.clone()));
#[cfg(feature = "tracing")]
tracing::Span::current().record("request_id", request_id.as_str());
let fut = self.inner.call(req);
Box::pin(async move {
let mut response = fut.await?;
response
.headers_mut()
.insert(REQUEST_ID_HEADER, header_value);
Ok(response)
})
}
}
pub fn get_request_id(headers: &http::HeaderMap) -> Option<String> {
headers
.get(REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[allow(clippy::type_complexity)]
fn echo_service() -> tower::util::ServiceFn<
impl Fn(
http::Request<()>,
) -> std::future::Ready<Result<http::Response<()>, std::convert::Infallible>>
+ Clone,
> {
tower::service_fn(|req: http::Request<()>| {
let observed_header = get_request_id(req.headers());
let observed_ext = req.extensions().get::<RequestId>().cloned();
let mut response = http::Response::new(());
if let Some(id) = observed_header {
response.headers_mut().insert(
"x-observed-request-id",
http::HeaderValue::from_str(&id).expect("uuid is a valid header value"),
);
}
if let Some(id) = observed_ext {
response.headers_mut().insert(
"x-observed-extension",
http::HeaderValue::from_str(&id.0).expect("uuid is a valid header value"),
);
}
std::future::ready(Ok::<_, std::convert::Infallible>(response))
})
}
#[tokio::test]
async fn test_generates_request_id_when_absent() {
let mut service = RequestIdLayer.layer(echo_service());
let response = Service::call(&mut service, http::Request::new(()))
.await
.expect("inner service is infallible");
let echoed = response
.headers()
.get(REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.expect("response echoes x-request-id");
assert!(uuid::Uuid::parse_str(echoed).is_ok(), "echoed id is a uuid");
assert_eq!(
response.headers().get("x-observed-request-id").unwrap(),
echoed,
"inner service saw the same id on the request"
);
assert_eq!(
response.headers().get("x-observed-extension").unwrap(),
echoed,
"inner service saw the same id in extensions"
);
}
#[tokio::test]
async fn test_preserves_incoming_request_id() {
let mut service = RequestIdLayer.layer(echo_service());
let mut request = http::Request::new(());
request.headers_mut().insert(
REQUEST_ID_HEADER,
http::HeaderValue::from_static("client-supplied-id"),
);
let response = Service::call(&mut service, request)
.await
.expect("inner service is infallible");
assert_eq!(
response.headers().get(REQUEST_ID_HEADER).unwrap(),
"client-supplied-id"
);
assert_eq!(
response.headers().get("x-observed-request-id").unwrap(),
"client-supplied-id"
);
}
#[tokio::test]
async fn test_replaces_oversized_request_id() {
let mut service = RequestIdLayer.layer(echo_service());
let mut request = http::Request::new(());
let oversized = "x".repeat(MAX_REQUEST_ID_LEN + 1);
request.headers_mut().insert(
REQUEST_ID_HEADER,
http::HeaderValue::from_str(&oversized).unwrap(),
);
let response = Service::call(&mut service, request)
.await
.expect("inner service is infallible");
let echoed = response
.headers()
.get(REQUEST_ID_HEADER)
.and_then(|v| v.to_str().ok())
.unwrap();
assert_ne!(echoed, oversized);
assert!(uuid::Uuid::parse_str(echoed).is_ok());
}
#[test]
fn test_get_request_id() {
let mut headers = http::HeaderMap::new();
headers.insert(REQUEST_ID_HEADER, http::HeaderValue::from_static("test-id"));
assert_eq!(get_request_id(&headers), Some("test-id".to_string()));
}
}