azure_core/http/policies/
client_request_id.rs1use 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#[derive(Debug)]
21pub struct ClientRequestIdPolicy(headers::HeaderName);
22
23impl ClientRequestIdPolicy {
24 pub const fn new() -> Self {
26 ClientRequestIdPolicy(headers::CLIENT_REQUEST_ID)
27 }
28
29 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 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 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 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 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 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 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 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 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 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 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 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 policy
228 .send(&ctx, &mut request, &[transport])
229 .await
230 .expect("Policy execution failed");
231 }
232}