Skip to main content

azure_identity/
azure_pipelines_credential.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use crate::{
5    env::Env, ClientAssertion, ClientAssertionCredential, ClientAssertionCredentialOptions,
6};
7use azure_core::{
8    credentials::{AccessToken, Secret, TokenCredential, TokenRequestOptions},
9    error::ErrorKind,
10    http::{
11        headers::{FromHeaders, HeaderName, Headers, AUTHORIZATION, CONTENT_LENGTH},
12        request::Request,
13        ClientMethodOptions, Method, Pipeline, PipelineSendOptions, StatusCode, Url,
14    },
15};
16use serde::Deserialize;
17use std::{any::type_name, borrow::Cow, convert::Infallible, fmt, sync::Arc};
18
19// cspell:ignore fedauthredirect msedge oidcrequesturi
20const OIDC_VARIABLE_NAME: &str = "SYSTEM_OIDCREQUESTURI";
21const OIDC_VERSION: &str = "7.1";
22const TFS_FEDAUTHREDIRECT_HEADER: HeaderName = HeaderName::from_static("x-tfs-fedauthredirect");
23
24const ALLOWED_HEADERS: &[&str] = &["x-msedge-ref", "x-vss-e2eid"];
25
26/// Authenticates an [Azure Pipelines service connection](https://learn.microsoft.com/azure/devops/pipelines/library/service-endpoints).
27pub struct AzurePipelinesCredential(ClientAssertionCredential<Client>);
28
29impl fmt::Debug for AzurePipelinesCredential {
30    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
31        f.debug_tuple(type_name::<Self>()).finish_non_exhaustive()
32    }
33}
34
35/// Options for constructing a new [`AzurePipelinesCredential`].
36#[derive(Default)]
37pub struct AzurePipelinesCredentialOptions {
38    /// Options for the [`ClientAssertionCredential`] used by the [`AzurePipelinesCredential`].
39    pub credential_options: ClientAssertionCredentialOptions,
40
41    #[cfg(test)]
42    pub(crate) env: Option<Env>,
43}
44
45impl fmt::Debug for AzurePipelinesCredentialOptions {
46    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
47        f.debug_struct(type_name::<Self>()).finish_non_exhaustive()
48    }
49}
50
51impl AzurePipelinesCredential {
52    /// Creates a new `AzurePipelinesCredential`.
53    ///
54    /// # Arguments
55    /// - `tenant_id`: The tenant (directory) ID of the service principal federated with the service connection.
56    /// - `client_id`: The client (application) ID of that service principal.
57    /// - `service_connection_id`: ID of the service connection to authenticate.
58    /// - `system_access_token`: Security token for the running build. See
59    ///   [Azure Pipelines documentation](https://learn.microsoft.com/azure/devops/pipelines/build/variables?view=azure-devops#systemaccesstoken)
60    ///   for an example showing how to get this value.
61    /// - `options`: Options for configuring the credential. If `None`, the credential uses its default options.
62    ///
63    pub fn new<T>(
64        tenant_id: String,
65        client_id: String,
66        service_connection_id: &str,
67        system_access_token: T,
68        options: Option<AzurePipelinesCredentialOptions>,
69    ) -> azure_core::Result<Arc<Self>>
70    where
71        T: Into<Secret>,
72    {
73        let system_access_token = system_access_token.into();
74
75        crate::validate_tenant_id(&tenant_id)?;
76        crate::validate_not_empty(&client_id, "no client ID specified")?;
77        crate::validate_not_empty(service_connection_id, "no service connection ID specified")?;
78        crate::validate_not_empty(
79            system_access_token.secret(),
80            "no system access token specified",
81        )?;
82
83        let mut options = options.unwrap_or_default();
84        options
85            .credential_options
86            .client_options
87            .logging
88            .additional_allowed_header_names
89            // the logging policy constructor will remove any duplicates
90            .extend(ALLOWED_HEADERS.iter().map(|&s| Cow::Borrowed(s)));
91
92        #[cfg(test)]
93        let env = options.env.unwrap_or_default();
94        #[cfg(not(test))]
95        let env = Env::default();
96
97        let endpoint = env
98            .var(OIDC_VARIABLE_NAME)
99            .map_err(|err| azure_core::Error::with_error(
100                ErrorKind::Credential,
101                err,
102                format!("no value for environment variable {OIDC_VARIABLE_NAME}. This should be set by Azure Pipelines"),
103            ))?;
104        let mut endpoint: Url = endpoint.parse().map_err(|err| {
105            azure_core::Error::with_error(
106                ErrorKind::Credential,
107                err,
108                format!("invalid URL for environment variable {OIDC_VARIABLE_NAME}"),
109            )
110        })?;
111        endpoint
112            .query_pairs_mut()
113            .append_pair("api-version", OIDC_VERSION)
114            .append_pair("serviceConnectionId", service_connection_id);
115        let pipeline = azure_core::http::Pipeline::new(
116            option_env!("CARGO_PKG_NAME"),
117            option_env!("CARGO_PKG_VERSION"),
118            options.credential_options.client_options.clone(),
119            Vec::default(),
120            Vec::default(),
121            None,
122        );
123        let client = Client {
124            endpoint,
125            pipeline: Arc::new(pipeline),
126            system_access_token,
127        };
128        let credential = ClientAssertionCredential::new_exclusive(
129            tenant_id,
130            client_id,
131            client,
132            stringify!(AzurePipelinesCredential),
133            Some(options.credential_options),
134        )?;
135
136        Ok(Arc::new(Self(credential)))
137    }
138}
139
140#[async_trait::async_trait]
141impl TokenCredential for AzurePipelinesCredential {
142    async fn get_token(
143        &self,
144        scopes: &[&str],
145        options: Option<TokenRequestOptions<'_>>,
146    ) -> azure_core::Result<AccessToken> {
147        self.0.get_token(scopes, options).await
148    }
149}
150
151#[derive(Debug)]
152struct Client {
153    endpoint: Url,
154    pipeline: Arc<Pipeline>,
155    system_access_token: Secret,
156}
157
158#[async_trait::async_trait]
159impl ClientAssertion for Client {
160    async fn secret(&self, options: Option<ClientMethodOptions<'_>>) -> azure_core::Result<String> {
161        let mut req = Request::new(self.endpoint.clone(), Method::Post);
162        req.insert_header(
163            AUTHORIZATION,
164            String::from("Bearer ") + self.system_access_token.secret(),
165        );
166        req.insert_header(TFS_FEDAUTHREDIRECT_HEADER, "Suppress");
167        req.insert_header(CONTENT_LENGTH, "0");
168
169        let options = options.unwrap_or_default();
170        let ctx = options.context.to_borrowed();
171        let resp = self
172            .pipeline
173            .send(
174                &ctx,
175                &mut req,
176                Some(PipelineSendOptions {
177                    skip_checks: true,
178                    ..Default::default()
179                }),
180            )
181            .await?;
182        let status = resp.status();
183        if status != StatusCode::Ok {
184            let err_headers: ErrorHeaders = resp.headers().get()?;
185            return Err(azure_core::Error::with_message(
186                ErrorKind::HttpResponse {
187                    status,
188                    error_code: Some(status.canonical_reason().to_string()),
189                    raw_response: Some(Box::new(resp)),
190                },
191                format!(
192                "{status} response from the OIDC endpoint. Check service connection ID and pipeline configuration. {err_headers}"
193            ),
194            ));
195        }
196
197        let assertion: Assertion = resp.into_body().json()?;
198        Ok(assertion.oidc_token.secret().to_string())
199    }
200}
201
202#[derive(Debug, Deserialize)]
203struct Assertion {
204    #[serde(rename = "oidcToken")]
205    oidc_token: Secret,
206}
207
208#[derive(Debug)]
209struct ErrorHeaders {
210    msedge_ref: Option<String>,
211    vss_e2eid: Option<String>,
212}
213
214const MSEDGE_REF: HeaderName = HeaderName::from_static("x-msedge-ref");
215const VSS_E2EID: HeaderName = HeaderName::from_static("x-vss-e2eid");
216
217impl FromHeaders for ErrorHeaders {
218    type Error = Infallible;
219
220    fn header_names() -> &'static [&'static str] {
221        ALLOWED_HEADERS
222    }
223
224    fn from_headers(headers: &Headers) -> Result<Option<Self>, Self::Error> {
225        Ok(Some(Self {
226            msedge_ref: headers.get_optional_string(&MSEDGE_REF),
227            vss_e2eid: headers.get_optional_string(&VSS_E2EID),
228        }))
229    }
230}
231
232impl fmt::Display for ErrorHeaders {
233    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
234        let mut v = f.debug_struct("Headers");
235        if let Some(ref msedge_ref) = self.msedge_ref {
236            v.field(MSEDGE_REF.as_str(), msedge_ref);
237        }
238        if let Some(ref vss_e2eid) = self.vss_e2eid {
239            v.field(VSS_E2EID.as_str(), vss_e2eid);
240        }
241        v.finish()
242    }
243}
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248    use crate::env::Env;
249    use azure_core::{
250        http::{AsyncRawResponse, ClientOptions, RawResponse, Transport},
251        Bytes,
252    };
253    use azure_core_test::http::MockHttpClient;
254    use futures::FutureExt as _;
255
256    #[test]
257    fn param_errors() {
258        assert!(AzurePipelinesCredential::new("".into(), "".into(), "", "", None).is_err());
259        assert!(AzurePipelinesCredential::new("_".into(), "".into(), "", "", None).is_err());
260        assert!(AzurePipelinesCredential::new("a".into(), "".into(), "", "", None).is_err());
261        assert!(AzurePipelinesCredential::new("a".into(), "b".into(), "", "", None).is_err());
262        assert!(AzurePipelinesCredential::new("a".into(), "b".into(), "c", "", None).is_err());
263
264        let options = AzurePipelinesCredentialOptions {
265            env: Some(Env::from(
266                &[(OIDC_VARIABLE_NAME, "http://localhost/get_token")][..],
267            )),
268            ..Default::default()
269        };
270        assert!(
271            AzurePipelinesCredential::new("a".into(), "b".into(), "c", "d", Some(options)).is_ok()
272        );
273    }
274
275    #[tokio::test]
276    async fn error_response() {
277        let expected_status = StatusCode::Forbidden;
278        let body = Bytes::from_static(b"content");
279        let mut headers = Headers::new();
280        headers.insert(MSEDGE_REF, "foo");
281        headers.insert(VSS_E2EID, "bar");
282        let expected_response =
283            RawResponse::from_bytes(expected_status, headers.clone(), body.clone());
284        let headers_for_mock = headers.clone();
285        let body_for_mock = body.clone();
286        let mock_client = MockHttpClient::new(move |req| {
287            assert_eq!(
288                req.url().as_str(),
289                "http://localhost/get_token?api-version=7.1&serviceConnectionId=c"
290            );
291            let headers = headers_for_mock.clone();
292            let body = body_for_mock.clone();
293
294            async move { Ok(AsyncRawResponse::from_bytes(expected_status, headers, body)) }.boxed()
295        });
296        let options = AzurePipelinesCredentialOptions {
297            credential_options: ClientAssertionCredentialOptions {
298                client_options: ClientOptions {
299                    transport: Some(Transport::new(Arc::new(mock_client))),
300                    ..Default::default()
301                },
302            },
303            env: Some(Env::from(
304                &[(OIDC_VARIABLE_NAME, "http://localhost/get_token")][..],
305            )),
306        };
307        let err = AzurePipelinesCredential::new("a".into(), "b".into(), "c", "d", Some(options))
308            .expect("credential")
309            .get_token(&["default"], None)
310            .await
311            .expect_err("expected error");
312
313        assert!(matches!(err.kind(), ErrorKind::Credential));
314        assert_eq!(
315            r#"AzurePipelinesCredential authentication failed. 403 response from the OIDC endpoint. Check service connection ID and pipeline configuration. Headers { x-msedge-ref: "foo", x-vss-e2eid: "bar" }
316To troubleshoot, visit https://aka.ms/azsdk/rust/identity/troubleshoot#apc"#,
317            err.to_string(),
318        );
319        match err
320            .downcast_ref::<azure_core::Error>()
321            .expect("returned error should wrap an azure_core::Error")
322            .kind()
323        {
324            ErrorKind::HttpResponse {
325                error_code: Some(reason),
326                raw_response: Some(response),
327                status,
328                ..
329            } => {
330                assert_eq!(status.canonical_reason(), reason.as_str());
331                assert_eq!(&expected_response, response.as_ref());
332                assert_eq!(expected_status, *status);
333            }
334            err => panic!("unexpected {:?}", err),
335        };
336    }
337
338    #[tokio::test]
339    async fn mock_request() {
340        let mock_client = MockHttpClient::new(|req| {
341            async move {
342                if req.url().as_str()
343                    == "http://localhost/get_token?api-version=7.1&serviceConnectionId=c"
344                {
345                    assert!(matches!(
346                        req.headers().get_str(&AUTHORIZATION),
347                        Ok(value) if value == "Bearer d",
348                    ));
349                    assert!(matches!(
350                        req.headers().get_str(&TFS_FEDAUTHREDIRECT_HEADER),
351                        Ok(value) if value == "Suppress",
352                    ));
353
354                    let mut headers = Headers::new();
355                    headers.insert(MSEDGE_REF, "foo");
356                    headers.insert(VSS_E2EID, "bar");
357
358                    return Ok(AsyncRawResponse::from_bytes(
359                        StatusCode::Ok,
360                        headers,
361                        Bytes::from_static(br#"{"oidcToken":"baz"}"#),
362                    ));
363                }
364
365                if req.url().as_str() == "https://login.microsoftonline.com/a/oauth2/v2.0/token" {
366                    return Ok(AsyncRawResponse::from_bytes(
367                        StatusCode::Ok,
368                        Headers::new(),
369                        Bytes::from_static(
370                            br#"{"token_type":"test","expires_in":0,"ext_expires_in":0,"access_token":"qux"}"#,
371                        ),
372                    ));
373                }
374
375                panic!("not supported")
376            }.boxed()
377        });
378        let options = AzurePipelinesCredentialOptions {
379            credential_options: ClientAssertionCredentialOptions {
380                client_options: ClientOptions {
381                    transport: Some(Transport::new(Arc::new(mock_client))),
382                    ..Default::default()
383                },
384            },
385            env: Some(Env::from(
386                &[(OIDC_VARIABLE_NAME, "http://localhost/get_token")][..],
387            )),
388        };
389        let credential =
390            AzurePipelinesCredential::new("a".into(), "b".into(), "c", "d", Some(options))
391                .expect("valid AzurePipelinesCredential");
392        let secret = credential
393            .get_token(&["default"], None)
394            .await
395            .expect("valid response");
396        assert_eq!(secret.token.secret(), "qux");
397    }
398}