1use 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#[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 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 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 pub fn access_key_id(&self) -> &str {
74 self.access_key_id.expose_secret()
75 }
76
77 pub fn secret_access_key(&self) -> &SecretString {
79 &self.secret_access_key
80 }
81
82 pub fn session_token(&self) -> Option<&SecretString> {
84 self.session_token.as_ref()
85 }
86
87 pub fn expires_at(&self) -> Option<OffsetDateTime> {
89 self.expires_at
90 }
91
92 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#[async_trait]
115pub trait CredentialsProvider: Send + Sync {
116 async fn provide_credentials(&self) -> Result<Credentials, S3Error>;
118}
119
120#[cfg(feature = "aws-credentials")]
127#[derive(Clone)]
128pub struct AwsDefaultCredentialsProvider {
129 inner: SharedCredentialsProvider,
130}
131
132#[cfg(feature = "aws-credentials")]
133impl AwsDefaultCredentialsProvider {
134 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 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#[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#[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#[derive(Clone)]
228pub struct StaticCredentialsProvider {
229 credentials: Credentials,
230}
231
232impl StaticCredentialsProvider {
233 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#[derive(Clone, Copy, Debug, Default)]
257pub struct EnvironmentCredentialsProvider;
258
259impl EnvironmentCredentialsProvider {
260 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
300pub struct CachedCredentialsProvider {
307 provider: Arc<dyn CredentialsProvider>,
308 refresh_before: Duration,
309 cached: Mutex<Option<Credentials>>,
310}
311
312impl CachedCredentialsProvider {
313 pub fn new(provider: Arc<dyn CredentialsProvider>) -> Self {
315 Self::with_refresh_before(provider, Duration::from_secs(60))
316 }
317
318 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 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}