1use super::{AZ_NAMESPACE_ATTRIBUTE, ERROR_TYPE_ATTRIBUTE};
5use crate::{
6 http::{Context, Request},
7 tracing::{Span, SpanKind, Tracer},
8};
9use ::tracing::trace;
10use std::{borrow::Cow, sync::Arc};
11use typespec_client_core::{
12 fmt::SafeDebug,
13 http::policies::{Policy, PolicyResult},
14 tracing::Attribute,
15};
16
17#[derive(SafeDebug, Clone)]
28pub struct PublicApiInstrumentationInformation {
29 #[safe(true)]
37 api_name: Cow<'static, str>,
38
39 attributes: Vec<Attribute>,
46}
47
48impl PublicApiInstrumentationInformation {
49 pub fn new(api_name: impl Into<Cow<'static, str>>, attributes: Vec<Attribute>) -> Self {
59 Self {
60 api_name: api_name.into(),
61 attributes,
62 }
63 }
64}
65
66#[derive(Clone, Debug)]
68pub(crate) struct PublicApiInstrumentationPolicy {
69 tracer: Option<Arc<dyn crate::tracing::Tracer>>,
70}
71
72impl PublicApiInstrumentationPolicy {
73 pub fn new(tracer: Option<Arc<dyn crate::tracing::Tracer>>) -> Self {
88 Self { tracer }
89 }
90}
91
92pub fn create_public_api_span(
113 ctx: &Context,
114 tracer: Option<Arc<dyn Tracer>>,
115 public_api_instrumentation: Option<PublicApiInstrumentationInformation>,
116) -> Option<Arc<dyn Span>> {
117 if ctx.value::<Arc<dyn Span>>().is_some() {
119 trace!(
120 "PublicApiPolicy: Nested call detected, forwarding request without instrumentation."
121 );
122 return None;
123 }
124
125 let info = public_api_instrumentation
128 .or_else(|| ctx.value::<PublicApiInstrumentationInformation>().cloned())?;
129
130 let tracer = match ctx.value::<Arc<dyn Tracer>>() {
132 Some(t) => t.clone(),
133 None => tracer?,
134 };
135
136 let mut span_attributes = info
140 .attributes
141 .iter()
142 .map(|attr| {
143 Attribute {
145 key: attr.key.clone(),
146 value: attr.value.clone(),
147 }
148 })
149 .collect::<Vec<_>>();
150
151 if let Some(namespace) = tracer.namespace() {
152 span_attributes.push(Attribute {
154 key: AZ_NAMESPACE_ATTRIBUTE.into(),
155 value: namespace.into(),
156 });
157 }
158
159 Some(tracer.start_span(info.api_name, SpanKind::Internal, span_attributes))
161}
162
163#[async_trait::async_trait]
164impl Policy for PublicApiInstrumentationPolicy {
165 async fn send(
166 &self,
167 ctx: &Context,
168 request: &mut Request,
169 next: &[Arc<dyn Policy>],
170 ) -> PolicyResult {
171 let Some(span) = create_public_api_span(ctx, self.tracer.clone(), None) else {
172 return next[0].send(ctx, request, &next[1..]).await;
173 };
174
175 let ctx = ctx.clone().with_value(span.clone());
177
178 let result = next[0].send(&ctx, request, &next[1..]).await;
179
180 if span.is_recording() {
182 match &result {
183 Err(e) => {
184 match e.kind() {
186 crate::error::ErrorKind::HttpResponse { status, .. } => {
187 span.set_attribute(ERROR_TYPE_ATTRIBUTE, status.to_string().into());
188
189 if status.is_server_error() {
192 span.set_status(crate::tracing::SpanStatus::Error {
193 description: "".to_string(),
194 });
195 }
196 }
197 _ => {
198 span.set_attribute(ERROR_TYPE_ATTRIBUTE, e.kind().to_string().into());
199 span.set_status(crate::tracing::SpanStatus::Error {
200 description: e.kind().to_string(),
201 });
202 }
203 }
204 }
205 Ok(response) => {
206 if response.status().is_server_error() {
209 span.set_status(crate::tracing::SpanStatus::Error {
210 description: "".to_string(),
211 });
212 }
213 if response.status().is_client_error() || response.status().is_server_error() {
214 span.set_attribute(
215 ERROR_TYPE_ATTRIBUTE,
216 response.status().to_string().into(),
217 );
218 }
219 }
220 }
221 }
222 span.end();
223 result
224 }
225}
226
227#[cfg(test)]
228mod tests {
229 use super::*;
231 use crate::{
232 http::{
233 headers::Headers,
234 policies::{create_public_api_span, RequestInstrumentationPolicy, TransportPolicy},
235 AsyncRawResponse, Method, StatusCode, Transport,
236 },
237 tracing::{SpanStatus, TracerProvider},
238 Result, Uuid,
239 };
240 use azure_core_test::{
241 http::MockHttpClient,
242 tracing::{
243 check_instrumentation_result, ExpectedSpanInformation, ExpectedTracerInformation,
244 MockTracingProvider,
245 },
246 };
247 use futures::future::BoxFuture;
248 use std::sync::Arc;
249 use typespec_client_core::http::DEFAULT_ALLOWED_QUERY_PARAMETERS;
250
251 async fn run_public_api_instrumentation_test<C>(
253 api_information: Option<PublicApiInstrumentationInformation>,
254 create_tracer: bool,
255 add_tracer_to_context: bool,
256 request: &mut Request,
257 callback: C,
258 ) -> Arc<MockTracingProvider>
259 where
260 C: FnMut(&Request) -> BoxFuture<'_, Result<AsyncRawResponse>> + Send + Sync + 'static,
261 {
262 let mock_tracer_provider = Arc::new(MockTracingProvider::new());
264
265 let tracer = if create_tracer {
266 Some(mock_tracer_provider.get_tracer(
267 add_tracer_to_context.then_some("test namespace"),
268 "test_crate",
269 Some("1.0.0"),
270 ))
271 } else {
272 None
273 };
274
275 let public_api_policy = {
276 let policy_tracer = tracer.clone();
277 Arc::new(PublicApiInstrumentationPolicy::new(policy_tracer))
278 };
279
280 let transport =
281 TransportPolicy::new(Transport::new(Arc::new(MockHttpClient::new(callback))));
282
283 let next: Vec<Arc<dyn Policy>> = vec![Arc::new(transport)];
284
285 let mut ctx = Context::default();
286 if let Some(t) = tracer {
287 if add_tracer_to_context {
288 ctx = ctx.with_value(t.clone());
290 }
291 }
292
293 if let Some(api_information) = api_information {
294 ctx = ctx.with_value(api_information);
296 }
297 let _result = public_api_policy.send(&ctx, request, &next).await;
298
299 mock_tracer_provider
300 }
301
302 async fn run_public_api_instrumentation_test_with_request_instrumentation<C>(
303 api_name: Option<&'static str>,
304 namespace: Option<&'static str>,
305 crate_name: Option<&'static str>,
306 version: Option<&'static str>,
307 request: &mut Request,
308 callback: C,
309 ) -> Arc<MockTracingProvider>
310 where
311 C: FnMut(&Request) -> BoxFuture<'_, Result<AsyncRawResponse>> + Send + Sync + 'static,
312 {
313 let mock_tracer_provider = Arc::new(MockTracingProvider::new());
314 let mock_tracer =
315 mock_tracer_provider.get_tracer(namespace, crate_name.unwrap_or("unknown"), version);
316
317 let public_api_policy = Arc::new(PublicApiInstrumentationPolicy::new(Some(
318 mock_tracer.clone(),
319 )));
320
321 let transport =
322 TransportPolicy::new(Transport::new(Arc::new(MockHttpClient::new(callback))));
323
324 let request_instrumentation_policy = RequestInstrumentationPolicy::new(
325 Some(mock_tracer.clone()),
326 (*DEFAULT_ALLOWED_QUERY_PARAMETERS).clone(),
327 );
328
329 let next: Vec<Arc<dyn Policy>> = vec![
330 Arc::new(request_instrumentation_policy),
331 Arc::new(transport),
332 ];
333 let public_api_information = PublicApiInstrumentationInformation::new(
334 api_name.unwrap_or("unknown"),
335 vec![Attribute {
336 key: "az.fake_attribute".into(),
337 value: "attribute value".into(),
338 }],
339 );
340
341 let ctx = Context::default()
343 .with_value(public_api_information)
344 .with_value(mock_tracer.clone());
345 let _result = public_api_policy.send(&ctx, request, &next).await;
346
347 mock_tracer_provider
348 }
349
350 #[test]
352 fn create_public_api_span_tests() {
353 let tracer =
354 Arc::new(MockTracingProvider::new()).get_tracer(Some("test"), "test", Some("1.0.0"));
355
356 {
358 let ctx = Context::default();
359 let span = create_public_api_span(&ctx, Some(tracer.clone()), None);
360 assert!(span.is_none(), "Should return None when no API info exists");
361 }
362 }
363
364 #[test]
366 fn create_public_api_span_tests_context_has_span() {
367 let tracer =
368 Arc::new(MockTracingProvider::new()).get_tracer(Some("test"), "test", Some("1.0.0"));
369 {
370 let existing_span = tracer.start_span("existing".into(), SpanKind::Internal, vec![]);
371 let ctx = Context::default().with_value(existing_span.clone());
372 let span = create_public_api_span(&ctx, Some(tracer.clone()), None);
373 assert!(
374 span.is_none(),
375 "Should return None when context already has a span"
376 );
377 }
378 }
379
380 #[test]
382 fn create_public_api_span_tests_public_api_information_from_param() {
383 let tracer =
384 Arc::new(MockTracingProvider::new()).get_tracer(Some("test"), "test", Some("1.0.0"));
385
386 {
388 let ctx = Context::default();
389 let span = create_public_api_span(
390 &ctx,
391 Some(tracer.clone()),
392 Some(PublicApiInstrumentationInformation::new(
393 "TestClient.test_api",
394 vec![],
395 )),
396 );
397 assert!(
398 span.is_some(),
399 "Should return Some when info exists as param"
400 );
401 }
402 }
403
404 #[test]
406 fn create_public_api_span_tests_public_api_info_no_tracer() {
407 {
408 let api_info = PublicApiInstrumentationInformation::new("TestClient.test_api", vec![]);
409 let ctx = Context::default().with_value(api_info);
410 let span = create_public_api_span(&ctx, None, None);
411 assert!(
412 span.is_none(),
413 "Should return None when no tracer is available"
414 );
415 }
416 }
417 #[test]
419 fn create_public_api_span_tests_api_info_and_tracer_from_context() {
420 let tracer =
421 Arc::new(MockTracingProvider::new()).get_tracer(Some("test"), "test", Some("1.0.0"));
422 {
423 let api_info = PublicApiInstrumentationInformation::new("TestClient.test_api", vec![]);
424 let ctx = Context::default()
425 .with_value(api_info)
426 .with_value(tracer.clone());
427 let span = create_public_api_span(&ctx, None, None);
428 assert!(
429 span.is_some(),
430 "Should create span when API info and tracer are available"
431 );
432 }
433 }
434 #[test]
436 fn create_public_api_span_tests_tracer_from_parameter() {
437 let tracer =
438 Arc::new(MockTracingProvider::new()).get_tracer(Some("test"), "test", Some("1.0.0"));
439 {
440 let api_info = PublicApiInstrumentationInformation::new(
441 "TestClient.test_api",
442 vec![Attribute {
443 key: "test.attribute".into(),
444 value: "test_value".into(),
445 }],
446 );
447 let ctx = Context::default().with_value(api_info);
448 let span = create_public_api_span(&ctx, Some(tracer.clone()), None);
449 assert!(span.is_some(), "Should create span with attributes");
450 }
451 }
452
453 #[tokio::test]
454 async fn public_api_instrumentation_no_public_api_info() {
455 let url = "http://example.com/path";
456 let mut request = Request::new(url.parse().unwrap(), Method::Get);
457
458 let mock_tracer = run_public_api_instrumentation_test(
459 None, true, true,
462 &mut request,
463 |req| {
464 Box::pin(async move {
465 assert_eq!(req.url().host_str(), Some("example.com"));
466 assert_eq!(req.method(), Method::Get);
467 Ok(AsyncRawResponse::from_bytes(
468 StatusCode::Ok,
469 Headers::new(),
470 vec![],
471 ))
472 })
473 },
474 )
475 .await;
476
477 check_instrumentation_result(
478 mock_tracer,
479 vec![ExpectedTracerInformation {
480 name: "test_crate",
481 version: Some("1.0.0"),
482 namespace: Some("test namespace"),
483 spans: vec![],
484 }],
485 );
486 }
487
488 #[tokio::test]
489 async fn public_api_instrumentation_no_tracer() {
490 let url = "http://example.com/path";
491 let mut request = Request::new(url.parse().unwrap(), Method::Get);
492
493 let mock_tracer = run_public_api_instrumentation_test(
494 Some(PublicApiInstrumentationInformation::new(
495 "MyClient.MyApi",
496 vec![],
497 )),
498 false, false, &mut request,
501 |req| {
502 Box::pin(async move {
503 assert_eq!(req.url().host_str(), Some("example.com"));
504 assert_eq!(req.method(), Method::Get);
505 Ok(AsyncRawResponse::from_bytes(
506 StatusCode::Ok,
507 Headers::new(),
508 vec![],
509 ))
510 })
511 },
512 )
513 .await;
514
515 check_instrumentation_result(mock_tracer, vec![]);
517 }
518
519 #[tokio::test]
520 async fn public_api_instrumentation_tracer_not_in_context() {
521 let url = "http://example.com/path";
522 let mut request = Request::new(url.parse().unwrap(), Method::Get);
523
524 let mock_tracer = run_public_api_instrumentation_test(
525 Some(PublicApiInstrumentationInformation::new(
526 "MyClient.MyApi",
527 vec![],
528 )),
529 true, false, &mut request,
532 |req| {
533 Box::pin(async move {
534 assert_eq!(req.url().host_str(), Some("example.com"));
535 assert_eq!(req.method(), Method::Get);
536 Ok(AsyncRawResponse::from_bytes(
537 StatusCode::Ok,
538 Headers::new(),
539 vec![],
540 ))
541 })
542 },
543 )
544 .await;
545
546 check_instrumentation_result(
547 mock_tracer,
548 vec![ExpectedTracerInformation {
549 name: "test_crate",
550 version: Some("1.0.0"),
551 namespace: None,
552 spans: vec![ExpectedSpanInformation {
553 span_name: "MyClient.MyApi",
554 status: SpanStatus::Unset,
555 kind: SpanKind::Internal,
556 span_id: Uuid::new_v4(),
557 ..Default::default()
558 }],
559 }],
560 )
561 }
562
563 #[tokio::test]
564 async fn simple_public_api_instrumentation_policy() {
565 let url = "http://example.com/path";
566 let mut request = Request::new(url.parse().unwrap(), Method::Get);
567
568 let mock_tracer = run_public_api_instrumentation_test(
569 Some(PublicApiInstrumentationInformation::new(
570 "MyClient.MyApi",
571 vec![],
572 )),
573 true, true,
575 &mut request,
576 |req| {
577 Box::pin(async move {
578 assert_eq!(req.url().host_str(), Some("example.com"));
579 assert_eq!(req.method(), Method::Get);
580 Ok(AsyncRawResponse::from_bytes(
581 StatusCode::Ok,
582 Headers::new(),
583 vec![],
584 ))
585 })
586 },
587 )
588 .await;
589
590 check_instrumentation_result(
591 mock_tracer,
592 vec![ExpectedTracerInformation {
593 name: "test_crate",
594 version: Some("1.0.0"),
595 namespace: Some("test namespace"),
596 spans: vec![ExpectedSpanInformation {
597 span_name: "MyClient.MyApi",
598 status: SpanStatus::Unset,
599 span_id: Uuid::new_v4(),
600 kind: SpanKind::Internal,
601 attributes: vec![(AZ_NAMESPACE_ATTRIBUTE, "test namespace".into())],
602 ..Default::default()
603 }],
604 }],
605 );
606 }
607
608 #[tokio::test]
609 async fn public_api_instrumentation_policy_with_error() {
610 let url = "http://example.com/path";
611 let mut request = Request::new(url.parse().unwrap(), Method::Get);
612
613 let mock_tracer = run_public_api_instrumentation_test(
614 Some(PublicApiInstrumentationInformation::new(
615 "MyClient.MyApi",
616 vec![],
617 )),
618 true,
619 true,
620 &mut request,
621 |req| {
622 Box::pin(async move {
623 assert_eq!(req.url().host_str(), Some("example.com"));
624 assert_eq!(req.method(), Method::Get);
625 Ok(AsyncRawResponse::from_bytes(
626 StatusCode::InternalServerError,
627 Headers::new(),
628 vec![],
629 ))
630 })
631 },
632 )
633 .await;
634
635 check_instrumentation_result(
636 mock_tracer.clone(),
637 vec![ExpectedTracerInformation {
638 name: "test_crate",
639 version: Some("1.0.0"),
640 namespace: Some("test namespace"),
641 spans: vec![ExpectedSpanInformation {
642 span_name: "MyClient.MyApi",
643 status: SpanStatus::Error {
644 description: "".to_string(),
645 },
646 kind: SpanKind::Internal,
647 span_id: Uuid::new_v4(),
648 parent_id: None,
649 attributes: vec![
650 (AZ_NAMESPACE_ATTRIBUTE, "test namespace".into()),
651 (ERROR_TYPE_ATTRIBUTE, "500".into()),
652 ],
653 ..Default::default()
654 }],
655 }],
656 );
657 }
658
659 #[tokio::test]
660 async fn public_api_instrumentation_policy_with_request_instrumentation() {
661 let url = "http://example.com/path_with_request";
662 let mut request = Request::new(url.parse().unwrap(), Method::Put);
663
664 let mock_tracer = run_public_api_instrumentation_test_with_request_instrumentation(
665 Some("MyClient.MyApi"),
666 Some("test.namespace"),
667 Some("test_crate"),
668 Some("1.0.0"),
669 &mut request,
670 |req| {
671 Box::pin(async move {
672 assert_eq!(req.url().host_str(), Some("example.com"));
673 assert_eq!(req.method(), Method::Put);
674 Ok(AsyncRawResponse::from_bytes(
675 StatusCode::Ok,
676 Headers::new(),
677 vec![],
678 ))
679 })
680 },
681 )
682 .await;
683
684 let parent_id = Uuid::new_v4();
685
686 check_instrumentation_result(
687 mock_tracer.clone(),
688 vec![ExpectedTracerInformation {
689 name: "test_crate",
690 version: Some("1.0.0"),
691 namespace: Some("test.namespace"),
692 spans: vec![
693 ExpectedSpanInformation {
694 span_name: "MyClient.MyApi",
695 status: SpanStatus::Unset,
696 kind: SpanKind::Internal,
697 span_id: parent_id,
698 parent_id: None,
699 attributes: vec![
700 (AZ_NAMESPACE_ATTRIBUTE, "test.namespace".into()),
701 ("az.fake_attribute", "attribute value".into()),
702 ],
703 ..Default::default()
704 },
705 ExpectedSpanInformation {
706 span_name: "PUT",
707 status: SpanStatus::Unset,
708 kind: SpanKind::Client,
709 span_id: Uuid::new_v4(),
710 parent_id: Some(parent_id),
711 attributes: vec![
712 (AZ_NAMESPACE_ATTRIBUTE, "test.namespace".into()),
713 ("http.request.method", "PUT".into()),
714 ("url.full", "http://example.com/path_with_request".into()),
715 ("server.address", "example.com".into()),
716 ("server.port", 80.into()),
717 ("http.response.status_code", 200.into()),
718 ],
719 ..Default::default()
720 },
721 ],
722 }],
723 );
724 }
725}