Skip to main content

azure_core/http/policies/
client_request_id.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use crate::{
5    http::{
6        headers::{self, Header as _},
7        policies::{Policy, PolicyResult},
8        request::options::ClientRequestId,
9        Context, Request,
10    },
11    Uuid,
12};
13use std::sync::Arc;
14
15/// Adds a `x-ms-client-request-id` (or custom) header to each request.
16///
17/// Clients can set a custom name by adding [`ClientRequestIdPolicy::with_header_name()`]
18/// to [`ClientOptions::per_call_policies`](crate::http::options::ClientOptions::per_call_policies).
19/// The default policy will not be added if a custom one has already been added.
20#[derive(Debug)]
21pub struct ClientRequestIdPolicy(headers::HeaderName);
22
23impl ClientRequestIdPolicy {
24    /// Creates a new policy using the default `x-ms-client-request-id` header.
25    pub const fn new() -> Self {
26        ClientRequestIdPolicy(headers::CLIENT_REQUEST_ID)
27    }
28
29    /// Creates a new policy using a custom header name.
30    ///
31    /// You can construct a new policy for a constant or static variable.
32    pub const fn with_header_name(header: &'static str) -> Self {
33        ClientRequestIdPolicy(headers::HeaderName::from_static(header))
34    }
35}
36
37impl Default for ClientRequestIdPolicy {
38    fn default() -> Self {
39        ClientRequestIdPolicy::new()
40    }
41}
42
43impl From<headers::HeaderName> for ClientRequestIdPolicy {
44    fn from(header_name: headers::HeaderName) -> Self {
45        Self(header_name)
46    }
47}
48
49#[async_trait::async_trait]
50impl Policy for ClientRequestIdPolicy {
51    async fn send(
52        &self,
53        ctx: &Context,
54        request: &mut Request,
55        next: &[Arc<dyn Policy>],
56    ) -> PolicyResult {
57        if request.headers().get_optional_str(&self.0).is_none() {
58            if let Some(request_id) = ctx.value::<ClientRequestId>() {
59                request.insert_header(self.0.clone(), request_id.value());
60            } else {
61                let request_id: String = Uuid::new_v4().into();
62                request.insert_header(self.0.clone(), request_id);
63            }
64        }
65
66        next[0].send(ctx, request, &next[1..]).await
67    }
68}
69
70#[cfg(test)]
71mod tests {
72    use super::*;
73    use crate::{
74        http::{headers, Method, Request, StatusCode},
75        Bytes,
76    };
77    use azure_core_test::http::MockHttpClient;
78    use futures::FutureExt;
79    use std::sync::Arc;
80    use typespec_client_core::http::{policies::TransportPolicy, AsyncRawResponse, Transport};
81
82    #[tokio::test]
83    async fn header_already_present() {
84        // Arrange
85        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
86        const EXISTING_REQUEST_ID: &str = "existing-request-id";
87        request.insert_header(headers::CLIENT_REQUEST_ID, EXISTING_REQUEST_ID);
88
89        let policy = ClientRequestIdPolicy::default();
90        let transport = Arc::new(MockHttpClient::new(|req| {
91            async move {
92                // Assert
93                let header_value = req
94                    .headers()
95                    .get_optional_str(&headers::CLIENT_REQUEST_ID)
96                    .expect("Header should be present");
97                assert_eq!(
98                    header_value, EXISTING_REQUEST_ID,
99                    "Header value should not change"
100                );
101
102                Ok(AsyncRawResponse::from_bytes(
103                    StatusCode::Ok,
104                    headers::Headers::new(),
105                    Bytes::new(),
106                ))
107            }
108            .boxed()
109        }));
110        let transport = Arc::new(TransportPolicy::new(Transport::new(transport)));
111        let ctx = Context::new();
112
113        // Act
114        policy
115            .send(&ctx, &mut request, &[transport])
116            .await
117            .expect("Policy execution failed");
118    }
119
120    #[tokio::test]
121    async fn header_not_present() {
122        // Arrange
123        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
124
125        let policy = ClientRequestIdPolicy::default();
126        let transport = Arc::new(MockHttpClient::new(|req| {
127            async move {
128                // Assert
129                let header_value = req
130                    .headers()
131                    .get_optional_str(&headers::CLIENT_REQUEST_ID)
132                    .expect("Header should be present");
133                assert!(!header_value.is_empty(), "Header value should be generated");
134
135                Ok(AsyncRawResponse::from_bytes(
136                    StatusCode::Ok,
137                    headers::Headers::new(),
138                    Bytes::new(),
139                ))
140            }
141            .boxed()
142        }));
143        let transport = Arc::new(TransportPolicy::new(Transport::new(transport)));
144        let ctx = Context::new();
145
146        // Act
147        policy
148            .send(&ctx, &mut request, &[transport])
149            .await
150            .expect("Policy execution failed");
151    }
152
153    #[tokio::test]
154    async fn custom_header_name_with_existing_value() {
155        // Arrange
156        let custom_header_name = headers::HeaderName::from_static("x-custom-request-id");
157        let existing_request_id = "custom-existing-request-id";
158
159        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
160        request.insert_header(custom_header_name.clone(), existing_request_id);
161
162        let policy = ClientRequestIdPolicy::with_header_name("x-custom-request-id");
163        let transport = Arc::new(MockHttpClient::new(move |req| {
164            let custom_header_name = custom_header_name.clone();
165            async move {
166                // Assert
167                let header_value = req
168                    .headers()
169                    .get_optional_str(&custom_header_name)
170                    .expect("Custom header should be present");
171                assert_eq!(
172                    header_value, existing_request_id,
173                    "Custom header value should not change"
174                );
175
176                Ok(AsyncRawResponse::from_bytes(
177                    StatusCode::Ok,
178                    headers::Headers::new(),
179                    Bytes::new(),
180                ))
181            }
182            .boxed()
183        }));
184        let transport = Arc::new(TransportPolicy::new(Transport::new(transport)));
185        let ctx = Context::new();
186
187        // Act
188        policy
189            .send(&ctx, &mut request, &[transport])
190            .await
191            .expect("Policy execution failed");
192    }
193
194    #[tokio::test]
195    async fn client_request_id_in_context() {
196        // Arrange
197        const CLIENT_REQUEST_ID: &str = "context-request-id";
198        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
199
200        let mut ctx = Context::new();
201        ctx.insert(ClientRequestId::new(CLIENT_REQUEST_ID.to_string()));
202
203        let policy = ClientRequestIdPolicy::default();
204        let transport = Arc::new(MockHttpClient::new(|req| {
205            async move {
206                // Assert
207                let header_value = req
208                    .headers()
209                    .get_optional_str(&headers::CLIENT_REQUEST_ID)
210                    .expect("Header should be present");
211                assert_eq!(
212                    header_value, CLIENT_REQUEST_ID,
213                    "Header value should match the client request ID from the context"
214                );
215
216                Ok(AsyncRawResponse::from_bytes(
217                    StatusCode::Ok,
218                    headers::Headers::new(),
219                    Bytes::new(),
220                ))
221            }
222            .boxed()
223        }));
224        let transport = Arc::new(TransportPolicy::new(Transport::new(transport)));
225
226        // Act
227        policy
228            .send(&ctx, &mut request, &[transport])
229            .await
230            .expect("Policy execution failed");
231    }
232}