Skip to main content

azure_core/http/
pipeline.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use super::policies::ClientRequestIdPolicy;
5use crate::{
6    error::CheckSuccessOptions,
7    http::{
8        check_success,
9        headers::{RETRY_AFTER_MS, X_MS_RETRY_AFTER_MS},
10        policies::{
11            Policy, PublicApiInstrumentationPolicy, RequestInstrumentationPolicy, UserAgentPolicy,
12        },
13        ClientOptions,
14    },
15};
16use std::{
17    any::{Any, TypeId},
18    sync::Arc,
19};
20use typespec_client_core::http::{
21    self, headers::RETRY_AFTER, policies::RetryHeaders, PipelineOptions,
22};
23
24/// Execution pipeline.
25///
26/// A pipeline follows a precise flow:
27///
28/// 1. Client library-specified per-call policies are executed. Per-call policies can fail and bail out of the pipeline
29///    immediately.
30/// 2. User-specified per-call policies in [`ClientOptions::per_call_policies`] are executed.
31/// 3. The retry policy is executed. It allows to re-execute the following policies.
32/// 4. Client library-specified per-retry policies. Per-retry polices are always executed at least once but are
33///    re-executed in case of retries.
34/// 5. User-specified per-retry policies in [`ClientOptions::per_try_policies`] are executed.
35/// 6. The transport policy is executed. Transport policy is always the last policy and is the policy that
36///    actually constructs the [`AsyncRawResponse`](http::AsyncRawResponse) to be passed up the pipeline.
37///
38/// A pipeline is immutable. In other words a policy can either succeed and call the following
39/// policy of fail and return to the calling policy. Arbitrary policy "skip" must be avoided (but
40/// cannot be enforced by code). All policies except Transport policy can assume there is another following policy (so
41/// `self.pipeline[0]` is always valid).
42#[derive(Debug, Clone)]
43pub struct Pipeline(http::Pipeline);
44
45/// Options for the [`Pipeline::send`] function.
46#[derive(Debug, Default)]
47pub struct PipelineSendOptions {
48    /// If true, skip all checks including [`check_success`].
49    pub skip_checks: bool,
50
51    /// Options for [`check_success`]. If `skip_checks` is true, this field is ignored.
52    pub check_success: CheckSuccessOptions,
53}
54
55/// Internal structure used to pass options to the core pipeline.
56#[derive(Debug, Default)]
57struct CorePipelineSendOptions {
58    check_success: CheckSuccessOptions,
59    skip_checks: bool,
60}
61
62impl PipelineSendOptions {
63    /// Deconstructs the `PipelineSendOptions` into its core components.
64    fn deconstruct(self) -> (CorePipelineSendOptions, Option<http::PipelineSendOptions>) {
65        (
66            CorePipelineSendOptions {
67                skip_checks: self.skip_checks,
68                check_success: self.check_success,
69            },
70            None,
71        )
72    }
73}
74
75/// Options for the [`Pipeline::stream`] function.
76#[derive(Debug, Default)]
77pub struct PipelineStreamOptions {
78    /// If true, skip all checks including [`check_success`].
79    pub skip_checks: bool,
80
81    /// Options for [`check_success`]. If `skip_checks` is true, this field is ignored.
82    pub check_success: CheckSuccessOptions,
83}
84
85/// Internal structure used to pass options to the core pipeline.
86#[derive(Debug, Default)]
87struct CorePipelineStreamOptions {
88    check_success: CheckSuccessOptions,
89    skip_checks: bool,
90}
91
92impl PipelineStreamOptions {
93    /// Deconstructs the `PipelineStreamOptions` into its core components.
94    fn deconstruct(
95        self,
96    ) -> (
97        CorePipelineStreamOptions,
98        Option<http::PipelineStreamOptions>,
99    ) {
100        (
101            CorePipelineStreamOptions {
102                skip_checks: self.skip_checks,
103                check_success: self.check_success,
104            },
105            None,
106        )
107    }
108}
109
110impl Pipeline {
111    /// Creates a new pipeline given the client library crate name and version,
112    /// alone with user-specified and client library-specified policies.
113    ///
114    /// Crates can simply pass `option_env!("CARGO_PKG_NAME")` and `option_env!("CARGO_PKG_VERSION")` for the
115    /// `crate_name` and `crate_version` arguments respectively.
116    ///
117    /// # Arguments
118    /// * `crate_name` - The name of the crate implementing the client library.
119    /// * `crate_version` - The version of the crate implementing the client library.
120    /// * `options` - The client options.
121    /// * `per_call_policies` - Policies to be executed per call, before the policies in `ClientOptions::per_call_policies`.
122    /// * `per_try_policies` - Policies to be executed per try, before the policies in `ClientOptions::per_try_policies`.
123    /// * `pipeline_options` - Additional options for the pipeline. If `None`, default options will be used.
124    pub fn new(
125        crate_name: Option<&'static str>,
126        crate_version: Option<&'static str>,
127        options: ClientOptions,
128        per_call_policies: Vec<Arc<dyn Policy>>,
129        per_try_policies: Vec<Arc<dyn Policy>>,
130        pipeline_options: Option<PipelineOptions>,
131    ) -> Self {
132        let (core_client_options, options) = options.deconstruct();
133
134        // Create a fallback tracer if no tracer provider is set.
135        // This is useful for service clients that have not yet been instrumented.
136        let tracer = core_client_options
137            .instrumentation
138            .tracer_provider
139            .map(|provider| {
140                // Note that the choice to use "None" as the namespace here
141                // is intentional.
142                // The `azure_namespace` parameter is used to populate the `az.namespace`
143                // span attribute, however that information is only known by the author of the
144                // client library, not the core library.
145                // It is also *not* a constant that can be derived from the crate information -
146                // it is a value that is determined from the list of resource providers
147                // listed [here](https://learn.microsoft.com/azure/azure-resource-manager/management/azure-services-resource-providers).
148                //
149                // This information can only come from the package owner. It doesn't make sense
150                // to burden all users of the azure_core pipeline with determining this
151                // information, so we use `None` here.
152                provider.get_tracer(None, crate_name.unwrap_or("Unknown"), crate_version)
153            });
154
155        let mut per_call_policies = per_call_policies.clone();
156        push_unique(&mut per_call_policies, ClientRequestIdPolicy::default());
157        if let Some(ref tracer) = tracer {
158            let public_api_policy = PublicApiInstrumentationPolicy::new(Some(tracer.clone()));
159            push_unique(&mut per_call_policies, public_api_policy);
160        }
161
162        let user_agent_policy =
163            UserAgentPolicy::new(crate_name, crate_version, &core_client_options.user_agent);
164        push_unique(&mut per_call_policies, user_agent_policy);
165
166        let mut per_try_policies = per_try_policies.clone();
167        if let Some(ref tracer) = tracer {
168            let request_instrumentation_policy = RequestInstrumentationPolicy::new(
169                Some(tracer.clone()),
170                core_client_options.allowed_query_params.clone(),
171            );
172            push_unique(&mut per_try_policies, request_instrumentation_policy);
173        }
174
175        let pipeline_options = pipeline_options.unwrap_or_else(|| PipelineOptions {
176            retry_headers: RetryHeaders {
177                retry_headers: vec![X_MS_RETRY_AFTER_MS, RETRY_AFTER_MS, RETRY_AFTER],
178            },
179            ..PipelineOptions::default()
180        });
181
182        Self(http::Pipeline::new(
183            options,
184            per_call_policies,
185            per_try_policies,
186            Some(pipeline_options),
187        ))
188    }
189
190    /// Sends a [`Request`](http::Request) through each configured [`Policy`] to get a [`RawResponse`](http::RawResponse) that is processed by each policy in reverse.
191    ///
192    /// # Arguments
193    /// * `ctx` - The context for the `Request`.
194    /// * `request` - The `Request` to send.
195    /// * `options` - Options for sending the `Request`, including check success options. If none, [`check_success`] will not be called.
196    ///
197    /// # Returns
198    ///
199    /// A [`http::RawResponse`] if the request was successful, or an `Error` if it failed.
200    /// If the response status code indicates an HTTP error, the function will attempt to parse the error response
201    /// body into an `ErrorResponse` and include it in the `Error`.
202    pub async fn send(
203        &self,
204        ctx: &http::Context<'_>,
205        request: &mut http::Request,
206        options: Option<PipelineSendOptions>,
207    ) -> crate::Result<http::RawResponse> {
208        let (core_send_options, send_options) = options.unwrap_or_default().deconstruct();
209        let result = self.0.send(ctx, request, send_options).await?;
210        if !core_send_options.skip_checks {
211            check_success(result, Some(core_send_options.check_success)).await
212        } else {
213            Ok(result)
214        }
215    }
216
217    /// Sends a [`Request`](http::Request) through each configured [`Policy`] to get a [`AsyncRawResponse`](http::AsyncRawResponse) that is processed by each policy in reverse.
218    ///
219    /// # Arguments
220    /// * `ctx` - The context for the `Request`.
221    /// * `request` - The `Request` to send.
222    /// * `options` - Options for sending the `Request`, including check success options. If none, [`check_success`] will not be called.
223    ///
224    /// # Returns
225    ///
226    /// A [`http::RawResponse`] if the request was successful, or an `Error` if it failed.
227    /// If the response status code indicates an HTTP error, the function will attempt to parse the error response
228    /// body into an `ErrorResponse` and include it in the `Error`.
229    pub async fn stream(
230        &self,
231        ctx: &http::Context<'_>,
232        request: &mut http::Request,
233        options: Option<PipelineStreamOptions>,
234    ) -> crate::Result<http::AsyncRawResponse> {
235        let (core_stream_options, stream_options) = options.unwrap_or_default().deconstruct();
236        let result = self.0.stream(ctx, request, stream_options).await?;
237        if !core_stream_options.skip_checks {
238            check_success(result, Some(core_stream_options.check_success)).await
239        } else {
240            Ok(result)
241        }
242    }
243}
244
245#[inline]
246fn push_unique<T: Policy + 'static>(policies: &mut Vec<Arc<dyn Policy>>, policy: T) {
247    if policies.iter().all(|p| TypeId::of::<T>() != p.type_id()) {
248        policies.push(Arc::new(policy));
249    }
250}
251
252#[cfg(test)]
253mod tests {
254    use super::*;
255    use crate::{
256        http::{
257            headers::{self, HeaderName, Headers},
258            policies::Policy,
259            request::options::ClientRequestId,
260            AsyncRawResponse, ClientOptions, Context, Method, Request, StatusCode, Transport,
261            UserAgentOptions,
262        },
263        Bytes,
264    };
265    use azure_core_test::http::MockHttpClient;
266    use futures::FutureExt as _;
267    use std::sync::Arc;
268
269    #[tokio::test]
270    async fn pipeline_with_custom_client_request_id_policy() {
271        // Arrange
272        const CUSTOM_HEADER_NAME: &str = "x-custom-request-id";
273        const CUSTOM_HEADER: HeaderName = HeaderName::from_static(CUSTOM_HEADER_NAME);
274        const CLIENT_REQUEST_ID: &str = "custom-request-id";
275
276        let mut ctx = Context::new();
277        ctx.insert(ClientRequestId::new(CLIENT_REQUEST_ID.to_string()));
278
279        let transport = Transport::new(Arc::new(MockHttpClient::new(|req| {
280            async {
281                // Assert
282                let header_value = req
283                    .headers()
284                    .get_optional_str(&CUSTOM_HEADER)
285                    .expect("Custom header should be present");
286                assert_eq!(
287                    header_value, CLIENT_REQUEST_ID,
288                    "Custom header value should match the client request ID"
289                );
290
291                Ok(AsyncRawResponse::from_bytes(
292                    StatusCode::Ok,
293                    Headers::new(),
294                    Bytes::new(),
295                ))
296            }
297            .boxed()
298        })));
299        let options = ClientOptions {
300            transport: Some(transport),
301            ..Default::default()
302        };
303
304        let per_call_policies: Vec<Arc<dyn Policy>> =
305            vec![
306                Arc::new(ClientRequestIdPolicy::with_header_name(CUSTOM_HEADER_NAME))
307                    as Arc<dyn Policy>,
308            ];
309        let per_retry_policies = vec![];
310
311        let pipeline = Pipeline::new(
312            Some("test-crate"),
313            Some("1.0.0"),
314            options,
315            per_call_policies,
316            per_retry_policies,
317            None,
318        );
319
320        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
321
322        // Act
323        pipeline
324            .send(&ctx, &mut request, None)
325            .await
326            .expect("Pipeline execution failed");
327    }
328
329    #[tokio::test]
330    async fn pipeline_without_client_request_id_policy() {
331        // Arrange
332        const CLIENT_REQUEST_ID: &str = "default-request-id";
333
334        let mut ctx = Context::new();
335        ctx.insert(ClientRequestId::new(CLIENT_REQUEST_ID.to_string()));
336
337        let transport = Transport::new(Arc::new(MockHttpClient::new(|req| {
338            async {
339                // Assert
340                let header_value = req
341                    .headers()
342                    .get_optional_str(&headers::CLIENT_REQUEST_ID)
343                    .expect("Default header should be present");
344                assert_eq!(
345                    header_value, CLIENT_REQUEST_ID,
346                    "Default header value should match the client request ID"
347                );
348
349                Ok(AsyncRawResponse::from_bytes(
350                    StatusCode::Ok,
351                    Headers::new(),
352                    Bytes::new(),
353                ))
354            }
355            .boxed()
356        })));
357        let options = ClientOptions {
358            transport: Some(transport),
359            ..Default::default()
360        };
361
362        let per_call_policies = vec![]; // No ClientRequestIdPolicy added
363        let per_retry_policies = vec![];
364
365        let pipeline = Pipeline::new(
366            Some("test-crate"),
367            Some("1.0.0"),
368            options,
369            per_call_policies,
370            per_retry_policies,
371            None,
372        );
373
374        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
375
376        // Act
377        pipeline
378            .send(&ctx, &mut request, None)
379            .await
380            .expect("Pipeline execution failed");
381    }
382
383    #[tokio::test]
384    async fn pipeline_with_user_agent_enabled_default() {
385        // Arrange
386        let ctx = Context::new();
387
388        let transport = Transport::new(Arc::new(MockHttpClient::new(|req| {
389            async {
390                // Assert
391                let user_agent = req
392                    .headers()
393                    .get_optional_str(&headers::USER_AGENT)
394                    .expect("User-Agent header should be present by default");
395                // The default user agent format is: azsdk-rust-<crate_name>/<crate_version> (<rustc_version>; <OS>; <ARCH>)
396                // Since we can't know the rustc version at runtime, just check the prefix and crate/version
397                assert!(
398                    user_agent.starts_with("azsdk-rust-test-crate/1.0.0 "),
399                    "User-Agent header should start with expected prefix, got: {}",
400                    user_agent
401                );
402
403                Ok(AsyncRawResponse::from_bytes(
404                    StatusCode::Ok,
405                    Headers::new(),
406                    Bytes::new(),
407                ))
408            }
409            .boxed()
410        })));
411        let options = ClientOptions {
412            transport: Some(transport),
413            ..Default::default()
414        };
415
416        let per_call_policies = vec![];
417        let per_retry_policies = vec![];
418
419        let pipeline = Pipeline::new(
420            Some("test-crate"),
421            Some("1.0.0"),
422            options,
423            per_call_policies,
424            per_retry_policies,
425            None,
426        );
427
428        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
429
430        // Act
431        pipeline
432            .send(&ctx, &mut request, None)
433            .await
434            .expect("Pipeline execution failed");
435    }
436
437    #[tokio::test]
438    async fn pipeline_with_custom_application_id() {
439        // Arrange
440        const CUSTOM_APPLICATION_ID: &str = "my-custom-app-2.1.0";
441        let ctx = Context::new();
442
443        let transport = Transport::new(Arc::new(MockHttpClient::new(|req| {
444            async {
445                // Assert
446                let user_agent = req
447                    .headers()
448                    .get_optional_str(&headers::USER_AGENT)
449                    .expect("User-Agent header should be present");
450                // The user agent should contain the custom application_id followed by the standard Azure SDK format
451                // Expected format: my-custom-app/2.1.0 azsdk-rust-test-crate/1.0.0 (<rustc_version>; <OS>; <ARCH>)
452                assert!(
453                    user_agent.starts_with("my-custom-app-2.1.0 azsdk-rust-test-crate/1.0.0 "),
454                    "User-Agent header should start with custom application_id and expected prefix, got: {}",
455                    user_agent
456                );
457
458                Ok(AsyncRawResponse::from_bytes(
459                    StatusCode::Ok,
460                    Headers::new(),
461                    Bytes::new(),
462                ))
463            }
464            .boxed()
465        })));
466
467        let user_agent_options = UserAgentOptions {
468            application_id: Some(CUSTOM_APPLICATION_ID.to_string()),
469        };
470
471        let options = ClientOptions {
472            transport: Some(transport),
473            user_agent: user_agent_options,
474            ..Default::default()
475        };
476
477        let per_call_policies = vec![];
478        let per_retry_policies = vec![];
479
480        let pipeline = Pipeline::new(
481            Some("test-crate"),
482            Some("1.0.0"),
483            options,
484            per_call_policies,
485            per_retry_policies,
486            None,
487        );
488
489        let mut request = Request::new("https://example.com".parse().unwrap(), Method::Get);
490
491        // Act
492        pipeline
493            .send(&ctx, &mut request, None)
494            .await
495            .expect("Pipeline execution failed");
496    }
497}