1use 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#[derive(Debug, Clone)]
43pub struct Pipeline(http::Pipeline);
44
45#[derive(Debug, Default)]
47pub struct PipelineSendOptions {
48 pub skip_checks: bool,
50
51 pub check_success: CheckSuccessOptions,
53}
54
55#[derive(Debug, Default)]
57struct CorePipelineSendOptions {
58 check_success: CheckSuccessOptions,
59 skip_checks: bool,
60}
61
62impl PipelineSendOptions {
63 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#[derive(Debug, Default)]
77pub struct PipelineStreamOptions {
78 pub skip_checks: bool,
80
81 pub check_success: CheckSuccessOptions,
83}
84
85#[derive(Debug, Default)]
87struct CorePipelineStreamOptions {
88 check_success: CheckSuccessOptions,
89 skip_checks: bool,
90}
91
92impl PipelineStreamOptions {
93 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 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 let tracer = core_client_options
137 .instrumentation
138 .tracer_provider
139 .map(|provider| {
140 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 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 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 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 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 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 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 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![]; 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 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 let ctx = Context::new();
387
388 let transport = Transport::new(Arc::new(MockHttpClient::new(|req| {
389 async {
390 let user_agent = req
392 .headers()
393 .get_optional_str(&headers::USER_AGENT)
394 .expect("User-Agent header should be present by default");
395 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 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 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 let user_agent = req
447 .headers()
448 .get_optional_str(&headers::USER_AGENT)
449 .expect("User-Agent header should be present");
450 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 pipeline
493 .send(&ctx, &mut request, None)
494 .await
495 .expect("Pipeline execution failed");
496 }
497}