Skip to main content

azure_core/
credentials.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4//! Azure authentication and authorization.
5
6use crate::Bytes;
7use serde::{Deserialize, Serialize};
8use std::{borrow::Cow, fmt};
9use typespec_client_core::{fmt::SafeDebug, http::ClientMethodOptions, time::OffsetDateTime};
10
11/// Represents a secret.
12///
13/// The [`Debug`] implementation will not print the secret.
14#[derive(Clone, Deserialize, Serialize, Eq)]
15pub struct Secret(Cow<'static, str>);
16
17impl Secret {
18    /// Create a new `Secret`.
19    pub fn new<T>(access_token: T) -> Self
20    where
21        T: Into<Cow<'static, str>>,
22    {
23        Self(access_token.into())
24    }
25
26    /// Get the secret value.
27    pub fn secret(&self) -> &str {
28        &self.0
29    }
30}
31
32// NOTE: this is a constant time compare, however LLVM may (and probably will)
33// optimize this in unexpected ways.
34impl PartialEq for Secret {
35    fn eq(&self, other: &Self) -> bool {
36        let a = self.secret();
37        let b = other.secret();
38
39        if a.len() != b.len() {
40            return false;
41        }
42
43        a.bytes()
44            .zip(b.bytes())
45            .fold(0, |acc, (a, b)| acc | (a ^ b))
46            == 0
47    }
48}
49
50impl From<String> for Secret {
51    fn from(access_token: String) -> Self {
52        Self::new(access_token)
53    }
54}
55
56impl From<&'static str> for Secret {
57    fn from(access_token: &'static str) -> Self {
58        Self::new(access_token)
59    }
60}
61
62impl fmt::Debug for Secret {
63    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
64        f.write_str("Secret")
65    }
66}
67
68/// Represents secret bytes, e.g., certificate data.
69///
70/// Neither the [`Debug`](fmt::Debug) nor the [`Display`](fmt::Display) implementation will print the bytes.
71#[derive(Clone, Eq)]
72pub struct SecretBytes(Vec<u8>);
73
74impl SecretBytes {
75    /// Create a new `SecretBytes`.
76    pub fn new(bytes: impl Into<Vec<u8>>) -> Self {
77        Self(bytes.into())
78    }
79
80    /// Get the secret bytes.
81    pub fn bytes(&self) -> &[u8] {
82        &self.0
83    }
84}
85
86// NOTE: this is a constant time compare, however LLVM may (and probably will)
87// optimize this in unexpected ways.
88impl PartialEq for SecretBytes {
89    fn eq(&self, other: &Self) -> bool {
90        let a = self.bytes();
91        let b = other.bytes();
92
93        if a.len() != b.len() {
94            return false;
95        }
96
97        a.iter().zip(b.iter()).fold(0, |acc, (a, b)| acc | (a ^ b)) == 0
98    }
99}
100
101impl fmt::Debug for SecretBytes {
102    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
103        f.write_str("SecretBytes")
104    }
105}
106
107impl fmt::Display for SecretBytes {
108    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
109        f.write_str("SecretBytes")
110    }
111}
112
113impl From<Bytes> for SecretBytes {
114    fn from(bytes: Bytes) -> Self {
115        Self(bytes.to_vec())
116    }
117}
118
119impl From<&[u8]> for SecretBytes {
120    fn from(bytes: &[u8]) -> Self {
121        Self(bytes.to_vec())
122    }
123}
124
125impl From<Vec<u8>> for SecretBytes {
126    fn from(bytes: Vec<u8>) -> Self {
127        Self(bytes)
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    #[test]
136    fn debug_does_not_print_bytes() {
137        let secret = SecretBytes::new(b"super-secret".to_vec());
138        assert_eq!("SecretBytes", format!("{secret:?}"));
139    }
140
141    #[test]
142    fn display_does_not_print_bytes() {
143        let secret = SecretBytes::new(b"super-secret".to_vec());
144        assert_eq!("SecretBytes", format!("{secret}"));
145    }
146
147    #[test]
148    fn eq_same_bytes() {
149        let a = SecretBytes::new(b"hello".to_vec());
150        let b = SecretBytes::new(b"hello".to_vec());
151        assert_eq!(a, b);
152    }
153
154    #[test]
155    fn ne_different_bytes() {
156        let a = SecretBytes::new(b"hello".to_vec());
157        let b = SecretBytes::new(b"world".to_vec());
158        assert_ne!(a, b);
159    }
160
161    #[test]
162    fn ne_different_lengths() {
163        let a = SecretBytes::new(b"hello".to_vec());
164        let b = SecretBytes::new(b"hello!".to_vec());
165        assert_ne!(a, b);
166    }
167
168    #[test]
169    fn from_bytes_type() {
170        let bytes = Bytes::from_static(b"data");
171        let secret = SecretBytes::from(bytes);
172        assert_eq!(b"data", secret.bytes());
173    }
174
175    #[test]
176    fn from_slice() {
177        let data: &[u8] = b"data";
178        let secret = SecretBytes::from(data);
179        assert_eq!(b"data", secret.bytes());
180    }
181
182    #[test]
183    fn from_vec() {
184        let secret = SecretBytes::from(b"data".to_vec());
185        assert_eq!(b"data", secret.bytes());
186    }
187}
188
189/// Represents an Azure service bearer access token with expiry information.
190#[derive(Debug, Clone, Serialize, Deserialize)]
191pub struct AccessToken {
192    /// Get the access token value.
193    pub token: Secret,
194    /// Gets the time when the provided token expires.
195    pub expires_on: OffsetDateTime,
196}
197
198impl AccessToken {
199    /// Create a new `AccessToken`.
200    pub fn new<T>(token: T, expires_on: OffsetDateTime) -> Self
201    where
202        T: Into<Secret>,
203    {
204        Self {
205            token: token.into(),
206            expires_on,
207        }
208    }
209}
210
211/// Options for getting a token from a [`TokenCredential`]
212#[derive(Clone, Default, SafeDebug)]
213pub struct TokenRequestOptions<'a> {
214    /// Method options to be used when requesting a token.
215    pub method_options: ClientMethodOptions<'a>,
216}
217
218/// Represents a credential that can acquire an Entra ID access token.
219///
220/// See the [azure_identity](https://docs.rs/azure_identity/latest/azure_identity/)
221/// crate for implementations.
222#[async_trait::async_trait]
223pub trait TokenCredential: Send + Sync + fmt::Debug {
224    /// Gets an [`AccessToken`] for the specified scopes
225    async fn get_token(
226        &self,
227        scopes: &[&str],
228        options: Option<TokenRequestOptions<'_>>,
229    ) -> crate::Result<AccessToken>;
230}