use std::sync::RwLock;
use std::time::{SystemTime, UNIX_EPOCH};
use alien_client_core::Result;
#[cfg(any(test, feature = "test-utils"))]
use alien_core::AwsServiceOverrides;
use alien_core::{AwsClientConfig, AwsCredentials};
use aws_credential_types::Credentials;
const REFRESH_BUFFER_SECS: u64 = 300;
const CREDENTIAL_LIFETIME_SECS: u64 = 3000;
#[derive(Debug, Clone)]
pub struct AwsCredentialProvider {
inner: std::sync::Arc<CredentialProviderInner>,
}
struct CredentialProviderInner {
original_config: AwsClientConfig,
cached: RwLock<CachedCredentials>,
refresh_lock: tokio::sync::Mutex<()>,
}
impl std::fmt::Debug for CredentialProviderInner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CredentialProviderInner")
.field("region", &self.original_config.region)
.field("account_id", &self.original_config.account_id)
.finish()
}
}
#[derive(Clone)]
struct CachedCredentials {
access_key_id: String,
secret_access_key: String,
session_token: Option<String>,
expires_at: u64,
}
impl AwsCredentialProvider {
pub async fn from_config(config: AwsClientConfig) -> Result<Self> {
let (cached, original_config) = match &config.credentials {
AwsCredentials::AccessKeys {
access_key_id,
secret_access_key,
session_token,
} => {
let cached = CachedCredentials {
access_key_id: access_key_id.clone(),
secret_access_key: secret_access_key.clone(),
session_token: session_token.clone(),
expires_at: 0, };
(cached, config)
}
AwsCredentials::WebIdentity { .. } => {
use crate::aws::AwsClientConfigExt;
let resolved = config.get_web_identity_credentials().await?;
let cached = match &resolved.credentials {
AwsCredentials::AccessKeys {
access_key_id,
secret_access_key,
session_token,
} => CachedCredentials {
access_key_id: access_key_id.clone(),
secret_access_key: secret_access_key.clone(),
session_token: session_token.clone(),
expires_at: now_secs() + CREDENTIAL_LIFETIME_SECS,
},
_ => unreachable!("get_web_identity_credentials always returns AccessKeys"),
};
(cached, config)
}
};
Ok(Self {
inner: std::sync::Arc::new(CredentialProviderInner {
original_config,
cached: RwLock::new(cached),
refresh_lock: tokio::sync::Mutex::new(()),
}),
})
}
#[cfg(any(test, feature = "test-utils"))]
pub fn from_config_sync(config: AwsClientConfig) -> Self {
let cached = match &config.credentials {
AwsCredentials::AccessKeys {
access_key_id,
secret_access_key,
session_token,
} => CachedCredentials {
access_key_id: access_key_id.clone(),
secret_access_key: secret_access_key.clone(),
session_token: session_token.clone(),
expires_at: 0,
},
AwsCredentials::WebIdentity { .. } => {
panic!("Cannot create sync credential provider from WebIdentity config")
}
};
Self {
inner: std::sync::Arc::new(CredentialProviderInner {
original_config: config,
cached: RwLock::new(cached),
refresh_lock: tokio::sync::Mutex::new(()),
}),
}
}
pub fn get_credentials(&self) -> Credentials {
let cached = self.inner.cached.read().unwrap();
Credentials::new(
cached.access_key_id.clone(),
cached.secret_access_key.clone(),
cached.session_token.clone(),
None,
"AwsCredentialProvider",
)
}
pub async fn ensure_fresh(&self) -> Result<()> {
let expires_at = {
let cached = self.inner.cached.read().unwrap();
cached.expires_at
};
if expires_at == 0 {
return Ok(());
}
let now = now_secs();
if now + REFRESH_BUFFER_SECS < expires_at {
return Ok(()); }
let _guard = self.inner.refresh_lock.lock().await;
{
let cached = self.inner.cached.read().unwrap();
if now_secs() + REFRESH_BUFFER_SECS < cached.expires_at {
return Ok(());
}
}
tracing::info!("Refreshing AWS credentials (IRSA token exchange)");
use crate::aws::AwsClientConfigExt;
let resolved = self
.inner
.original_config
.get_web_identity_credentials()
.await?;
match &resolved.credentials {
AwsCredentials::AccessKeys {
access_key_id,
secret_access_key,
session_token,
} => {
let mut cached = self.inner.cached.write().unwrap();
cached.access_key_id = access_key_id.clone();
cached.secret_access_key = secret_access_key.clone();
cached.session_token = session_token.clone();
cached.expires_at = now_secs() + CREDENTIAL_LIFETIME_SECS;
}
_ => unreachable!("get_web_identity_credentials always returns AccessKeys"),
}
tracing::info!("AWS credentials refreshed successfully");
Ok(())
}
pub fn region(&self) -> &str {
&self.inner.original_config.region
}
pub fn account_id(&self) -> &str {
&self.inner.original_config.account_id
}
pub fn get_service_endpoint_option(&self, service_name: &str) -> Option<&str> {
self.inner
.original_config
.service_overrides
.as_ref()
.and_then(|overrides| overrides.endpoints.get(service_name))
.map(|s| s.as_str())
}
pub fn get_service_endpoint(&self, service_name: &str, default_endpoint: &str) -> String {
self.get_service_endpoint_option(service_name)
.map(|s| s.to_string())
.unwrap_or_else(|| default_endpoint.to_string())
}
pub fn config(&self) -> &AwsClientConfig {
&self.inner.original_config
}
pub async fn with_region(&self, region: &str) -> Result<Self> {
let mut config = self.inner.original_config.clone();
config.region = region.to_string();
Self::from_config(config).await
}
#[cfg(any(test, feature = "test-utils"))]
pub fn with_service_overrides(self, overrides: AwsServiceOverrides) -> Self {
let mut config = self.inner.original_config.clone();
config.service_overrides = Some(overrides);
let cached = self.inner.cached.read().unwrap().clone();
Self {
inner: std::sync::Arc::new(CredentialProviderInner {
original_config: config,
cached: RwLock::new(cached),
refresh_lock: tokio::sync::Mutex::new(()),
}),
}
}
}
fn now_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_secs()
}