Skip to main content

rest/
resilience.rs

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/// Propagates a client request ID or assigns one when it is absent.
16#[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    /// Retrieves the request ID assigned by this middleware.
39    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/// The request ID attached to an Actix request extension.
48#[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}