Skip to main content

s3_wire/credentials/
mod.rs

1//! Credential values and providers.
2
3use std::env;
4use std::fmt;
5use std::sync::Arc;
6use std::time::Duration;
7
8use async_trait::async_trait;
9use secrecy::{ExposeSecret, SecretString};
10use time::OffsetDateTime;
11use tokio::sync::Mutex;
12
13use crate::error::{ErrorCategory, RetryClassification, S3Error};
14
15/// AWS-compatible credentials and optional expiration.
16///
17/// Secret values are zeroized on drop. Its `Debug` implementation redacts all
18/// credential material, including the access-key identifier.
19#[derive(Clone)]
20pub struct Credentials {
21    access_key_id: SecretString,
22    secret_access_key: SecretString,
23    session_token: Option<SecretString>,
24    expires_at: Option<OffsetDateTime>,
25}
26
27impl Credentials {
28    /// Creates non-expiring credentials.
29    pub fn new(
30        access_key_id: impl Into<String>,
31        secret_access_key: impl Into<String>,
32        session_token: Option<String>,
33    ) -> Result<Self, S3Error> {
34        Self::with_expiration(access_key_id, secret_access_key, session_token, None)
35    }
36
37    /// Creates credentials with an optional expiration time.
38    pub fn with_expiration(
39        access_key_id: impl Into<String>,
40        secret_access_key: impl Into<String>,
41        session_token: Option<String>,
42        expires_at: Option<OffsetDateTime>,
43    ) -> Result<Self, S3Error> {
44        let access_key_id = access_key_id.into();
45        let secret_access_key = secret_access_key.into();
46        if access_key_id.is_empty() {
47            return Err(S3Error::configuration(
48                "credential access key identifier must not be empty",
49            ));
50        }
51        if secret_access_key.is_empty() {
52            return Err(S3Error::configuration(
53                "credential secret access key must not be empty",
54            ));
55        }
56        if session_token.as_deref() == Some("") {
57            return Err(S3Error::configuration(
58                "credential session token must not be empty when present",
59            ));
60        }
61        Ok(Self {
62            access_key_id: access_key_id.into(),
63            secret_access_key: secret_access_key.into(),
64            session_token: session_token.map(Into::into),
65            expires_at,
66        })
67    }
68
69    /// Returns the access-key identifier for signing.
70    pub fn access_key_id(&self) -> &str {
71        self.access_key_id.expose_secret()
72    }
73
74    /// Returns the secret access key for signing.
75    pub fn secret_access_key(&self) -> &SecretString {
76        &self.secret_access_key
77    }
78
79    /// Returns the optional session token for signing.
80    pub fn session_token(&self) -> Option<&SecretString> {
81        self.session_token.as_ref()
82    }
83
84    /// Returns the credential expiration time, when supplied by the provider.
85    pub fn expires_at(&self) -> Option<OffsetDateTime> {
86        self.expires_at
87    }
88
89    /// Returns whether the credentials expire at or before `time`.
90    pub fn expires_by(&self, time: OffsetDateTime) -> bool {
91        self.expires_at.is_some_and(|expiration| expiration <= time)
92    }
93}
94
95impl fmt::Debug for Credentials {
96    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
97        formatter
98            .debug_struct("Credentials")
99            .field("access_key_id", &"[REDACTED]")
100            .field("secret_access_key", &"[REDACTED]")
101            .field(
102                "session_token",
103                &self.session_token.as_ref().map(|_| "[REDACTED]"),
104            )
105            .field("expires_at", &self.expires_at)
106            .finish()
107    }
108}
109
110/// Asynchronous source of AWS-compatible credentials.
111#[async_trait]
112pub trait CredentialsProvider: Send + Sync {
113    /// Resolves credentials for a request.
114    async fn provide_credentials(&self) -> Result<Credentials, S3Error>;
115}
116
117/// Provider backed by an immutable credential value.
118#[derive(Clone)]
119pub struct StaticCredentialsProvider {
120    credentials: Credentials,
121}
122
123impl StaticCredentialsProvider {
124    /// Creates a static provider.
125    pub fn new(credentials: Credentials) -> Self {
126        Self { credentials }
127    }
128}
129
130impl fmt::Debug for StaticCredentialsProvider {
131    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
132        formatter
133            .debug_struct("StaticCredentialsProvider")
134            .field("credentials", &"[REDACTED]")
135            .finish()
136    }
137}
138
139#[async_trait]
140impl CredentialsProvider for StaticCredentialsProvider {
141    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
142        Ok(self.credentials.clone())
143    }
144}
145
146/// Provider for the conventional AWS credential environment variables.
147#[derive(Clone, Copy, Debug, Default)]
148pub struct EnvironmentCredentialsProvider;
149
150impl EnvironmentCredentialsProvider {
151    /// Creates an environment credential provider.
152    pub const fn new() -> Self {
153        Self
154    }
155}
156
157#[async_trait]
158impl CredentialsProvider for EnvironmentCredentialsProvider {
159    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
160        let access_key_id = read_required_env("AWS_ACCESS_KEY_ID")?;
161        let secret_access_key = read_required_env("AWS_SECRET_ACCESS_KEY")?;
162        let session_token = read_optional_env("AWS_SESSION_TOKEN")?;
163        Credentials::new(access_key_id, secret_access_key, session_token)
164    }
165}
166
167fn read_required_env(name: &'static str) -> Result<String, S3Error> {
168    read_optional_env(name)?.ok_or_else(|| {
169        S3Error::new(
170            ErrorCategory::Authentication,
171            format!("required credential environment variable {name} is not set"),
172            RetryClassification::Never,
173        )
174    })
175}
176
177fn read_optional_env(name: &'static str) -> Result<Option<String>, S3Error> {
178    env::var_os(name)
179        .map(|value| {
180            value.into_string().map_err(|_| {
181                S3Error::new(
182                    ErrorCategory::Authentication,
183                    format!("credential environment variable {name} is not valid Unicode"),
184                    RetryClassification::Never,
185                )
186            })
187        })
188        .transpose()
189}
190
191/// Expiration-aware provider that coalesces concurrent refresh requests.
192///
193/// A refresh is performed while holding a Tokio mutex. Concurrent callers wait for
194/// that single refresh and then consume the same cached result, preventing refresh
195/// storms. If an early refresh fails, credentials that have not actually expired are
196/// returned until a later call can retry.
197pub struct CachedCredentialsProvider {
198    provider: Arc<dyn CredentialsProvider>,
199    refresh_before: Duration,
200    cached: Mutex<Option<Credentials>>,
201}
202
203impl CachedCredentialsProvider {
204    /// Wraps a provider and refreshes expiring credentials one minute early.
205    pub fn new(provider: Arc<dyn CredentialsProvider>) -> Self {
206        Self::with_refresh_before(provider, Duration::from_secs(60))
207    }
208
209    /// Wraps a provider with a custom early-refresh window.
210    pub fn with_refresh_before(
211        provider: Arc<dyn CredentialsProvider>,
212        refresh_before: Duration,
213    ) -> Self {
214        Self {
215            provider,
216            refresh_before,
217            cached: Mutex::new(None),
218        }
219    }
220
221    fn refresh_deadline(&self, now: OffsetDateTime) -> OffsetDateTime {
222        let refresh_before =
223            time::Duration::try_from(self.refresh_before).unwrap_or(time::Duration::MAX);
224        now.saturating_add(refresh_before)
225    }
226
227    /// Clears the current cache entry so the next request refreshes it.
228    pub async fn invalidate(&self) {
229        *self.cached.lock().await = None;
230    }
231}
232
233impl fmt::Debug for CachedCredentialsProvider {
234    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
235        formatter
236            .debug_struct("CachedCredentialsProvider")
237            .field("provider", &"[REDACTED]")
238            .field("refresh_before", &self.refresh_before)
239            .field("cached", &"[REDACTED]")
240            .finish()
241    }
242}
243
244#[async_trait]
245impl CredentialsProvider for CachedCredentialsProvider {
246    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
247        let mut cached = self.cached.lock().await;
248        let now = OffsetDateTime::now_utc();
249        if let Some(credentials) = cached.as_ref()
250            && !credentials.expires_by(self.refresh_deadline(now))
251        {
252            return Ok(credentials.clone());
253        }
254
255        match self.provider.provide_credentials().await {
256            Ok(credentials) => {
257                *cached = Some(credentials.clone());
258                Ok(credentials)
259            }
260            Err(refresh_error) => {
261                if let Some(credentials) = cached.as_ref()
262                    && !credentials.expires_by(now)
263                {
264                    return Ok(credentials.clone());
265                }
266                Err(refresh_error)
267            }
268        }
269    }
270}
271
272#[cfg(test)]
273mod tests {
274    use std::sync::atomic::{AtomicUsize, Ordering};
275
276    use futures_util::future::join_all;
277    use proptest::prelude::*;
278
279    use super::*;
280
281    struct CountingProvider {
282        calls: AtomicUsize,
283    }
284
285    #[async_trait]
286    impl CredentialsProvider for CountingProvider {
287        async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
288            self.calls.fetch_add(1, Ordering::SeqCst);
289            tokio::task::yield_now().await;
290            Credentials::new(
291                "visible-access-id",
292                "visible-secret",
293                Some("visible-token".into()),
294            )
295        }
296    }
297
298    #[test]
299    fn debug_redacts_all_credential_material() {
300        let credentials = Credentials::new(
301            "visible-access-id",
302            "visible-secret",
303            Some("visible-token".into()),
304        )
305        .unwrap();
306        let debug = format!("{credentials:?}");
307        assert!(!debug.contains("visible-access-id"));
308        assert!(!debug.contains("visible-secret"));
309        assert!(!debug.contains("visible-token"));
310    }
311
312    proptest! {
313        #[test]
314        fn arbitrary_credential_values_are_redacted(
315            suffix in "[A-Za-z0-9]{12,32}"
316        ) {
317            let access = format!("ACCESS-{suffix}");
318            let secret = format!("SECRET-{suffix}");
319            let token = format!("TOKEN-{suffix}");
320            let credentials = Credentials::new(
321                access.clone(),
322                secret.clone(),
323                Some(token.clone()),
324            ).unwrap();
325            let debug = format!("{credentials:?}");
326            prop_assert!(!debug.contains(&access));
327            prop_assert!(!debug.contains(&secret));
328            prop_assert!(!debug.contains(&token));
329        }
330    }
331
332    #[tokio::test]
333    async fn concurrent_cache_misses_share_one_refresh() {
334        let provider = Arc::new(CountingProvider {
335            calls: AtomicUsize::new(0),
336        });
337        let cached = Arc::new(CachedCredentialsProvider::new(provider.clone()));
338        let requests = (0..32).map(|_| {
339            let cached = cached.clone();
340            async move { cached.provide_credentials().await.unwrap() }
341        });
342
343        let results = join_all(requests).await;
344        assert_eq!(results.len(), 32);
345        assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
346    }
347}