1use actix_web::{
2 body::MessageBody,
3 dev::{Service, ServiceRequest, ServiceResponse, Transform},
4 http::header::{HeaderName, HeaderValue},
5 Error, HttpMessage, HttpRequest,
6};
7use futures::future::{ok, LocalBoxFuture, Ready};
8use std::{
9 sync::atomic::{AtomicU64, Ordering},
10 task::{Context, Poll},
11};
12
13static NEXT_REQUEST_ID: AtomicU64 = AtomicU64::new(1);
14
15#[derive(Clone)]
17pub struct RequestId {
18 header: HeaderName,
19}
20
21impl Default for RequestId {
22 fn default() -> Self {
23 Self {
24 header: HeaderName::from_static("x-request-id"),
25 }
26 }
27}
28
29impl RequestId {
30 pub fn new() -> Self {
31 Self::default()
32 }
33
34 pub fn with_header(header: HeaderName) -> Self {
35 Self { header }
36 }
37
38 pub fn request_id(request: &HttpRequest) -> Option<String> {
40 request
41 .extensions()
42 .get::<RequestIdValue>()
43 .map(|value| value.0.clone())
44 }
45}
46
47#[derive(Debug, Clone, PartialEq, Eq)]
49pub struct RequestIdValue(String);
50
51impl RequestIdValue {
52 pub fn as_str(&self) -> &str {
53 &self.0
54 }
55}
56
57impl<S, B> Transform<S, ServiceRequest> for RequestId
58where
59 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
60 S::Future: 'static,
61 B: MessageBody + 'static,
62{
63 type Response = ServiceResponse<B>;
64 type Error = Error;
65 type Transform = RequestIdMiddleware<S>;
66 type InitError = ();
67 type Future = Ready<Result<Self::Transform, Self::InitError>>;
68
69 fn new_transform(&self, service: S) -> Self::Future {
70 ok(RequestIdMiddleware {
71 service,
72 header: self.header.clone(),
73 })
74 }
75}
76
77pub struct RequestIdMiddleware<S> {
78 service: S,
79 header: HeaderName,
80}
81
82impl<S, B> Service<ServiceRequest> for RequestIdMiddleware<S>
83where
84 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
85 S::Future: 'static,
86 B: MessageBody + 'static,
87{
88 type Response = ServiceResponse<B>;
89 type Error = Error;
90 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
91
92 fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
93 self.service.poll_ready(context)
94 }
95
96 fn call(&self, request: ServiceRequest) -> Self::Future {
97 let request_id = request
98 .headers()
99 .get(&self.header)
100 .and_then(|value| value.to_str().ok())
101 .filter(|value| !value.is_empty())
102 .map(ToOwned::to_owned)
103 .unwrap_or_else(|| format!("{:016x}", NEXT_REQUEST_ID.fetch_add(1, Ordering::Relaxed)));
104 request
105 .extensions_mut()
106 .insert(RequestIdValue(request_id.clone()));
107 let header = self.header.clone();
108 let future = self.service.call(request);
109
110 Box::pin(async move {
111 let mut response = future.await?;
112 response.headers_mut().insert(
113 header,
114 HeaderValue::from_str(&request_id)
115 .expect("request IDs are sourced from validated HTTP headers"),
116 );
117 Ok(response)
118 })
119 }
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125 use actix_web::{test, web, App, HttpResponse};
126
127 #[actix_rt::test]
128 async fn preserves_the_client_request_id() {
129 let app = test::init_service(App::new().wrap(RequestId::new()).route(
130 "/",
131 web::get().to(|request: HttpRequest| async move {
132 HttpResponse::Ok().body(RequestId::request_id(&request).unwrap())
133 }),
134 ))
135 .await;
136
137 let response = test::call_service(
138 &app,
139 test::TestRequest::get()
140 .uri("/")
141 .insert_header(("x-request-id", "from-client"))
142 .to_request(),
143 )
144 .await;
145
146 assert_eq!(
147 response.headers().get("x-request-id").unwrap(),
148 "from-client"
149 );
150 }
151}