Skip to main content

azure_core/http/policies/instrumentation/
public_api_instrumentation.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use 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/// Information about the public API being instrumented.
18///
19/// This struct is used to pass information about the public API being instrumented
20/// to the `PublicApiInstrumentationPolicy`.
21///
22/// It contains the name of the API, which is used to create a span for distributed tracing
23/// and any additional per-API attributes that might be needed for instrumentation.
24///
25/// If the `PublicApiInstrumentationPolicy` policy detects a `PublicApiInstrumentationInformation` in the context,
26/// it will create a span with the API name and any additional attributes.
27#[derive(SafeDebug, Clone)]
28pub struct PublicApiInstrumentationInformation {
29    /// The name of the API being instrumented.
30    ///
31    /// The API name should be in the form of `<client>.<api>`, where
32    /// `<client>` is the name of the service client and `<api>` is the name of the API.
33    ///
34    /// For example, if the service client is `MyClient` and the API is `my_api`,
35    /// the API name should be `MyClient.my_api`.
36    #[safe(true)]
37    api_name: Cow<'static, str>,
38
39    /// Additional attributes to be added to the span for this API.
40    ///
41    /// These attributes can provide additional information about the API being instrumented.
42    /// See [Library-specific attributes](https://github.com/Azure/azure-sdk/blob/main/docs/tracing/distributed-tracing-conventions.md#library-specific-attributes)
43    /// for more information.
44    ///
45    attributes: Vec<Attribute>,
46}
47
48impl PublicApiInstrumentationInformation {
49    /// Creates a new `PublicApiInstrumentationInformation`.
50    ///
51    /// # Arguments
52    /// - `api_name`: The name of the API being instrumented.
53    /// - `attributes`: Additional attributes to be added to the span for this API.
54    ///
55    /// # Returns
56    /// A new instance of `PublicApiInstrumentationInformation`.
57    ///
58    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/// Sets distributed tracing information for HTTP requests.
67#[derive(Clone, Debug)]
68pub(crate) struct PublicApiInstrumentationPolicy {
69    tracer: Option<Arc<dyn crate::tracing::Tracer>>,
70}
71
72impl PublicApiInstrumentationPolicy {
73    /// Creates a new `PublicApiInstrumentationPolicy`.
74    ///
75    ///
76    /// # Returns
77    /// A new instance of `PublicApiInstrumentationPolicy`.
78    ///
79    /// # Note
80    /// This policy will only create a tracer if a tracing provider is provided in the options.
81    ///
82    /// This policy will create a tracer that can be used to instrument HTTP requests.
83    /// However this tracer is only used when the client method is NOT instrumented.
84    /// A part of the client method instrumentation sets a client-specific tracer into the
85    /// request `[Context]` which will be used instead of the tracer from this policy.
86    ///
87    pub fn new(tracer: Option<Arc<dyn crate::tracing::Tracer>>) -> Self {
88        Self { tracer }
89    }
90}
91
92/// Creates a span for the public API instrumentation policy.
93///
94/// This function creates a span for the public API instrumentation policy based on the
95/// public API information in the context.
96///
97/// If no PublicApiInstrumentationInformation is provided, then this function will look in the `Context`
98/// for a `PublicApiInstrumentationInformation` value, if it is not present, it will return `None`.
99///
100/// # Arguments
101/// - `ctx`: The context containing the public API information.
102/// - `tracer`: An optional tracer to use for creating the span.
103/// - `public_api_instrumentation`: Optional public API instrumentation information.
104///
105/// # Returns
106/// An optional span if the public API information is present and a tracer is available.
107///
108/// If the context already has a span, it will return `None` to avoid nested spans.
109/// If the context does not have a tracer it will use the value of the `tracer` argument.
110/// If no tracer can be determined, it will return `None`.
111///
112pub 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 there is a span in the context, we're a nested call, so we just want to forward the request.
118    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    // We next confirm if the context has public API instrumentation information.
126    // Without a public API information, we skip instrumentation.
127    let info = public_api_instrumentation
128        .or_else(|| ctx.value::<PublicApiInstrumentationInformation>().cloned())?;
129
130    // Get the tracer from either the context or the policy.
131    let tracer = match ctx.value::<Arc<dyn Tracer>>() {
132        Some(t) => t.clone(),
133        None => tracer?,
134    };
135
136    // We now have public API information and a tracer.
137    // Calculate the span attributes based on the public API information and
138    // tracer.
139    let mut span_attributes = info
140        .attributes
141        .iter()
142        .map(|attr| {
143            // Convert the attribute to a span attribute.
144            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        // If the tracer has a namespace, we set it as an attribute.
153        span_attributes.push(Attribute {
154            key: AZ_NAMESPACE_ATTRIBUTE.into(),
155            value: namespace.into(),
156        });
157    }
158
159    // Create a span with the public API information and attributes.
160    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        // Now add the span to the context, so that it can be used by the next policies.
176        let ctx = ctx.clone().with_value(span.clone());
177
178        let result = next[0].send(&ctx, request, &next[1..]).await;
179
180        // Don't bother setting attributes if the span isn't recording.
181        if span.is_recording() {
182            match &result {
183                Err(e) => {
184                    // If the request failed, we set the error type on the span.
185                    match e.kind() {
186                        crate::error::ErrorKind::HttpResponse { status, .. } => {
187                            span.set_attribute(ERROR_TYPE_ATTRIBUTE, status.to_string().into());
188
189                            // 5xx status codes SHOULD set status to Error.
190                            // The description should not be set because it can be inferred from "http.response.status_code".
191                            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                    // 5xx status codes SHOULD set status to Error.
207                    // The description should not be set because it can be inferred from "http.response.status_code".
208                    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    // cspell: ignore traceparent
230    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    // Test just the public API instrumentation policy without request instrumentation.
252    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        // Add the public API information and tracer to the context so that it can be used by the policy.
263        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                // If we have a tracer, add it to the context.
289                ctx = ctx.with_value(t.clone());
290            }
291        }
292
293        if let Some(api_information) = api_information {
294            // If we have public API information, add it to the context.
295            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        // Add the public API information and tracer to the context so that it can be used by the policy.
342        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    // Tests for the create_public_api_span function.
351    #[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        // Test when context has no PublicApiInstrumentationInformation
357        {
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 when context already has a span
365    #[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    // Tests for the create_public_api_span function.
381    #[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        // Test when context has no PublicApiInstrumentationInformation
387        {
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 with API info but no tracer
405    #[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 with API info and tracer from context
418    #[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 with API info, tracer from parameter, and attributes
435    #[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, // No public API information.
460            true, // Create tracer.
461            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, // Create tracer.
499            false, // Add tracer to context.
500            &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        // No tracer should be created, so we expect no spans.
516        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,  // Create tracer.
530            false, // Add tracer to context.
531            &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, // Create tracer.
574            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}