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#[cfg(feature = "aws-credentials")]
16use aws_credential_types::provider::{ProvideCredentials as _, SharedCredentialsProvider};
17
18/// AWS-compatible credentials and optional expiration.
19///
20/// Secret values are zeroized on drop. Its `Debug` implementation redacts all
21/// credential material, including the access-key identifier.
22#[derive(Clone)]
23pub struct Credentials {
24    access_key_id: SecretString,
25    secret_access_key: SecretString,
26    session_token: Option<SecretString>,
27    expires_at: Option<OffsetDateTime>,
28}
29
30impl Credentials {
31    /// Creates non-expiring credentials.
32    pub fn new(
33        access_key_id: impl Into<String>,
34        secret_access_key: impl Into<String>,
35        session_token: Option<String>,
36    ) -> Result<Self, S3Error> {
37        Self::with_expiration(access_key_id, secret_access_key, session_token, None)
38    }
39
40    /// Creates credentials with an optional expiration time.
41    pub fn with_expiration(
42        access_key_id: impl Into<String>,
43        secret_access_key: impl Into<String>,
44        session_token: Option<String>,
45        expires_at: Option<OffsetDateTime>,
46    ) -> Result<Self, S3Error> {
47        let access_key_id = access_key_id.into();
48        let secret_access_key = secret_access_key.into();
49        if access_key_id.is_empty() {
50            return Err(S3Error::configuration(
51                "credential access key identifier must not be empty",
52            ));
53        }
54        if secret_access_key.is_empty() {
55            return Err(S3Error::configuration(
56                "credential secret access key must not be empty",
57            ));
58        }
59        if session_token.as_deref() == Some("") {
60            return Err(S3Error::configuration(
61                "credential session token must not be empty when present",
62            ));
63        }
64        Ok(Self {
65            access_key_id: access_key_id.into(),
66            secret_access_key: secret_access_key.into(),
67            session_token: session_token.map(Into::into),
68            expires_at,
69        })
70    }
71
72    /// Returns the access-key identifier for signing.
73    pub fn access_key_id(&self) -> &str {
74        self.access_key_id.expose_secret()
75    }
76
77    /// Returns the secret access key for signing.
78    pub fn secret_access_key(&self) -> &SecretString {
79        &self.secret_access_key
80    }
81
82    /// Returns the optional session token for signing.
83    pub fn session_token(&self) -> Option<&SecretString> {
84        self.session_token.as_ref()
85    }
86
87    /// Returns the credential expiration time, when supplied by the provider.
88    pub fn expires_at(&self) -> Option<OffsetDateTime> {
89        self.expires_at
90    }
91
92    /// Returns whether the credentials expire at or before `time`.
93    pub fn expires_by(&self, time: OffsetDateTime) -> bool {
94        self.expires_at.is_some_and(|expiration| expiration <= time)
95    }
96}
97
98impl fmt::Debug for Credentials {
99    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
100        formatter
101            .debug_struct("Credentials")
102            .field("access_key_id", &"[REDACTED]")
103            .field("secret_access_key", &"[REDACTED]")
104            .field(
105                "session_token",
106                &self.session_token.as_ref().map(|_| "[REDACTED]"),
107            )
108            .field("expires_at", &self.expires_at)
109            .finish()
110    }
111}
112
113/// Asynchronous source of AWS-compatible credentials.
114#[async_trait]
115pub trait CredentialsProvider: Send + Sync {
116    /// Resolves credentials for a request.
117    async fn provide_credentials(&self) -> Result<Credentials, S3Error>;
118}
119
120/// Adapter for AWS's standard renewable credential provider chain.
121///
122/// This provider is available only with the `aws-credentials` feature so the
123/// default `s3-wire` dependency graph remains small. The chain covers shared
124/// profiles, environment credentials, web identity, ECS container credentials,
125/// EC2 IMDSv2, credential processes, and assume-role profiles.
126#[cfg(feature = "aws-credentials")]
127#[derive(Clone)]
128pub struct AwsDefaultCredentialsProvider {
129    inner: SharedCredentialsProvider,
130}
131
132#[cfg(feature = "aws-credentials")]
133impl AwsDefaultCredentialsProvider {
134    /// Builds the standard AWS credential chain.
135    pub async fn new() -> Self {
136        let chain = aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
137            .build()
138            .await;
139        Self {
140            inner: SharedCredentialsProvider::new(chain),
141        }
142    }
143
144    /// Builds the standard AWS credential chain for a named shared profile.
145    pub async fn for_profile(profile: &str) -> Result<Self, S3Error> {
146        if profile.is_empty() {
147            return Err(S3Error::configuration("AWS profile name must not be empty"));
148        }
149        let chain = aws_config::default_provider::credentials::DefaultCredentialsChain::builder()
150            .profile_name(profile)
151            .build()
152            .await;
153        Ok(Self {
154            inner: SharedCredentialsProvider::new(chain),
155        })
156    }
157}
158
159#[cfg(feature = "aws-credentials")]
160impl fmt::Debug for AwsDefaultCredentialsProvider {
161    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
162        formatter
163            .debug_struct("AwsDefaultCredentialsProvider")
164            .field("inner", &"[REDACTED]")
165            .finish()
166    }
167}
168
169#[cfg(feature = "aws-credentials")]
170#[async_trait]
171impl CredentialsProvider for AwsDefaultCredentialsProvider {
172    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
173        let credentials = self.inner.provide_credentials().await.map_err(|error| {
174            S3Error::new(
175                ErrorCategory::Authentication,
176                "AWS default credential chain did not resolve credentials",
177                RetryClassification::Never,
178            )
179            .with_source(error)
180        })?;
181        Credentials::with_expiration(
182            credentials.access_key_id(),
183            credentials.secret_access_key(),
184            credentials.session_token().map(ToOwned::to_owned),
185            credentials.expiry().map(OffsetDateTime::from),
186        )
187    }
188}
189
190/// Resolves a signing region through AWS's standard environment, shared
191/// profile, and IMDS region chain.
192///
193/// # Errors
194///
195/// Returns an error when the chain does not resolve a region.
196#[cfg(feature = "aws-credentials")]
197pub async fn resolve_default_aws_region() -> Result<String, S3Error> {
198    aws_config::default_provider::region::DefaultRegionChain::builder()
199        .build()
200        .region()
201        .await
202        .map(|region| region.as_ref().to_owned())
203        .ok_or_else(|| S3Error::configuration("AWS default region chain did not resolve a region"))
204}
205
206/// Resolves a signing region through AWS's standard chain for a named profile.
207///
208/// # Errors
209///
210/// Returns an error for an empty profile or when the chain does not resolve a
211/// region.
212#[cfg(feature = "aws-credentials")]
213pub async fn resolve_aws_region_for_profile(profile: &str) -> Result<String, S3Error> {
214    if profile.is_empty() {
215        return Err(S3Error::configuration("AWS profile name must not be empty"));
216    }
217    aws_config::default_provider::region::DefaultRegionChain::builder()
218        .profile_name(profile)
219        .build()
220        .region()
221        .await
222        .map(|region| region.as_ref().to_owned())
223        .ok_or_else(|| S3Error::configuration("AWS profile did not resolve a region"))
224}
225
226/// Provider backed by an immutable credential value.
227#[derive(Clone)]
228pub struct StaticCredentialsProvider {
229    credentials: Credentials,
230}
231
232impl StaticCredentialsProvider {
233    /// Creates a static provider.
234    pub fn new(credentials: Credentials) -> Self {
235        Self { credentials }
236    }
237}
238
239impl fmt::Debug for StaticCredentialsProvider {
240    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
241        formatter
242            .debug_struct("StaticCredentialsProvider")
243            .field("credentials", &"[REDACTED]")
244            .finish()
245    }
246}
247
248#[async_trait]
249impl CredentialsProvider for StaticCredentialsProvider {
250    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
251        Ok(self.credentials.clone())
252    }
253}
254
255/// Provider for the conventional AWS credential environment variables.
256#[derive(Clone, Copy, Debug, Default)]
257pub struct EnvironmentCredentialsProvider;
258
259impl EnvironmentCredentialsProvider {
260    /// Creates an environment credential provider.
261    pub const fn new() -> Self {
262        Self
263    }
264}
265
266#[async_trait]
267impl CredentialsProvider for EnvironmentCredentialsProvider {
268    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
269        let access_key_id = read_required_env("AWS_ACCESS_KEY_ID")?;
270        let secret_access_key = read_required_env("AWS_SECRET_ACCESS_KEY")?;
271        let session_token = read_optional_env("AWS_SESSION_TOKEN")?;
272        Credentials::new(access_key_id, secret_access_key, session_token)
273    }
274}
275
276fn read_required_env(name: &'static str) -> Result<String, S3Error> {
277    read_optional_env(name)?.ok_or_else(|| {
278        S3Error::new(
279            ErrorCategory::Authentication,
280            format!("required credential environment variable {name} is not set"),
281            RetryClassification::Never,
282        )
283    })
284}
285
286fn read_optional_env(name: &'static str) -> Result<Option<String>, S3Error> {
287    env::var_os(name)
288        .map(|value| {
289            value.into_string().map_err(|_| {
290                S3Error::new(
291                    ErrorCategory::Authentication,
292                    format!("credential environment variable {name} is not valid Unicode"),
293                    RetryClassification::Never,
294                )
295            })
296        })
297        .transpose()
298}
299
300/// Expiration-aware provider that coalesces concurrent refresh requests.
301///
302/// A refresh is performed while holding a Tokio mutex. Concurrent callers wait for
303/// that single refresh and then consume the same cached result, preventing refresh
304/// storms. If an early refresh fails, credentials that have not actually expired are
305/// returned until a later call can retry.
306pub struct CachedCredentialsProvider {
307    provider: Arc<dyn CredentialsProvider>,
308    refresh_before: Duration,
309    cached: Mutex<Option<Credentials>>,
310}
311
312impl CachedCredentialsProvider {
313    /// Wraps a provider and refreshes expiring credentials one minute early.
314    pub fn new(provider: Arc<dyn CredentialsProvider>) -> Self {
315        Self::with_refresh_before(provider, Duration::from_secs(60))
316    }
317
318    /// Wraps a provider with a custom early-refresh window.
319    pub fn with_refresh_before(
320        provider: Arc<dyn CredentialsProvider>,
321        refresh_before: Duration,
322    ) -> Self {
323        Self {
324            provider,
325            refresh_before,
326            cached: Mutex::new(None),
327        }
328    }
329
330    fn refresh_deadline(&self, now: OffsetDateTime) -> OffsetDateTime {
331        let refresh_before =
332            time::Duration::try_from(self.refresh_before).unwrap_or(time::Duration::MAX);
333        now.saturating_add(refresh_before)
334    }
335
336    /// Clears the current cache entry so the next request refreshes it.
337    pub async fn invalidate(&self) {
338        *self.cached.lock().await = None;
339    }
340}
341
342impl fmt::Debug for CachedCredentialsProvider {
343    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
344        formatter
345            .debug_struct("CachedCredentialsProvider")
346            .field("provider", &"[REDACTED]")
347            .field("refresh_before", &self.refresh_before)
348            .field("cached", &"[REDACTED]")
349            .finish()
350    }
351}
352
353#[async_trait]
354impl CredentialsProvider for CachedCredentialsProvider {
355    async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
356        let mut cached = self.cached.lock().await;
357        let now = OffsetDateTime::now_utc();
358        if let Some(credentials) = cached.as_ref()
359            && !credentials.expires_by(self.refresh_deadline(now))
360        {
361            return Ok(credentials.clone());
362        }
363
364        match self.provider.provide_credentials().await {
365            Ok(credentials) => {
366                *cached = Some(credentials.clone());
367                Ok(credentials)
368            }
369            Err(refresh_error) => {
370                if let Some(credentials) = cached.as_ref()
371                    && !credentials.expires_by(now)
372                {
373                    return Ok(credentials.clone());
374                }
375                Err(refresh_error)
376            }
377        }
378    }
379}
380
381#[cfg(test)]
382mod tests {
383    use std::sync::atomic::{AtomicUsize, Ordering};
384
385    use futures_util::future::join_all;
386    use proptest::prelude::*;
387
388    use super::*;
389
390    struct CountingProvider {
391        calls: AtomicUsize,
392    }
393
394    #[async_trait]
395    impl CredentialsProvider for CountingProvider {
396        async fn provide_credentials(&self) -> Result<Credentials, S3Error> {
397            self.calls.fetch_add(1, Ordering::SeqCst);
398            tokio::task::yield_now().await;
399            Credentials::new(
400                "visible-access-id",
401                "visible-secret",
402                Some("visible-token".into()),
403            )
404        }
405    }
406
407    #[test]
408    fn debug_redacts_all_credential_material() {
409        let credentials = Credentials::new(
410            "visible-access-id",
411            "visible-secret",
412            Some("visible-token".into()),
413        )
414        .unwrap();
415        let debug = format!("{credentials:?}");
416        assert!(!debug.contains("visible-access-id"));
417        assert!(!debug.contains("visible-secret"));
418        assert!(!debug.contains("visible-token"));
419    }
420
421    proptest! {
422        #[test]
423        fn arbitrary_credential_values_are_redacted(
424            suffix in "[A-Za-z0-9]{12,32}"
425        ) {
426            let access = format!("ACCESS-{suffix}");
427            let secret = format!("SECRET-{suffix}");
428            let token = format!("TOKEN-{suffix}");
429            let credentials = Credentials::new(
430                access.clone(),
431                secret.clone(),
432                Some(token.clone()),
433            ).unwrap();
434            let debug = format!("{credentials:?}");
435            prop_assert!(!debug.contains(&access));
436            prop_assert!(!debug.contains(&secret));
437            prop_assert!(!debug.contains(&token));
438        }
439    }
440
441    #[tokio::test]
442    async fn concurrent_cache_misses_share_one_refresh() {
443        let provider = Arc::new(CountingProvider {
444            calls: AtomicUsize::new(0),
445        });
446        let cached = Arc::new(CachedCredentialsProvider::new(provider.clone()));
447        let requests = (0..32).map(|_| {
448            let cached = cached.clone();
449            async move { cached.provide_credentials().await.unwrap() }
450        });
451
452        let results = join_all(requests).await;
453        assert_eq!(results.len(), 32);
454        assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
455    }
456}