s3_wire/credentials/
mod.rs1use 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#[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 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 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 pub fn access_key_id(&self) -> &str {
71 self.access_key_id.expose_secret()
72 }
73
74 pub fn secret_access_key(&self) -> &SecretString {
76 &self.secret_access_key
77 }
78
79 pub fn session_token(&self) -> Option<&SecretString> {
81 self.session_token.as_ref()
82 }
83
84 pub fn expires_at(&self) -> Option<OffsetDateTime> {
86 self.expires_at
87 }
88
89 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#[async_trait]
112pub trait CredentialsProvider: Send + Sync {
113 async fn provide_credentials(&self) -> Result<Credentials, S3Error>;
115}
116
117#[derive(Clone)]
119pub struct StaticCredentialsProvider {
120 credentials: Credentials,
121}
122
123impl StaticCredentialsProvider {
124 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#[derive(Clone, Copy, Debug, Default)]
148pub struct EnvironmentCredentialsProvider;
149
150impl EnvironmentCredentialsProvider {
151 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
191pub struct CachedCredentialsProvider {
198 provider: Arc<dyn CredentialsProvider>,
199 refresh_before: Duration,
200 cached: Mutex<Option<Credentials>>,
201}
202
203impl CachedCredentialsProvider {
204 pub fn new(provider: Arc<dyn CredentialsProvider>) -> Self {
206 Self::with_refresh_before(provider, Duration::from_secs(60))
207 }
208
209 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 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}