use std::sync::Arc;
use rskit_errors::AppResult;
use rskit_util::{SecretString, env};
use super::{SigningConfig, TransportAuth};
pub const DEFAULT_TOKEN_USERNAME: &str = "x-access-token";
pub trait AuthProvider: Send + Sync {
fn transport_auth(&self, remote: Option<&str>) -> AppResult<Option<TransportAuth>>;
fn signing_config(&self) -> AppResult<Option<SigningConfig>> {
Ok(None)
}
}
impl AuthProvider for Arc<dyn AuthProvider> {
fn transport_auth(&self, remote: Option<&str>) -> AppResult<Option<TransportAuth>> {
(**self).transport_auth(remote)
}
fn signing_config(&self) -> AppResult<Option<SigningConfig>> {
(**self).signing_config()
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct DefaultAuthProvider;
impl AuthProvider for DefaultAuthProvider {
fn transport_auth(&self, _remote: Option<&str>) -> AppResult<Option<TransportAuth>> {
Ok(None)
}
}
#[derive(Debug, Clone)]
pub struct StaticAuthProvider {
auth: TransportAuth,
}
impl StaticAuthProvider {
#[must_use]
pub fn new(auth: TransportAuth) -> Self {
Self { auth }
}
}
impl AuthProvider for StaticAuthProvider {
fn transport_auth(&self, _remote: Option<&str>) -> AppResult<Option<TransportAuth>> {
Ok(Some(self.auth.clone()))
}
}
#[derive(Debug, Clone)]
pub struct EnvTokenAuthProvider {
vars: Vec<String>,
username: String,
}
impl EnvTokenAuthProvider {
#[must_use]
pub fn with_var(name: impl Into<String>) -> Self {
Self {
vars: vec![name.into()],
username: DEFAULT_TOKEN_USERNAME.to_string(),
}
}
#[must_use]
pub fn with_vars<I, S>(names: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
Self {
vars: names.into_iter().map(Into::into).collect(),
username: DEFAULT_TOKEN_USERNAME.to_string(),
}
}
#[must_use]
pub fn with_username(mut self, username: impl Into<String>) -> Self {
self.username = username.into();
self
}
}
impl AuthProvider for EnvTokenAuthProvider {
fn transport_auth(&self, _remote: Option<&str>) -> AppResult<Option<TransportAuth>> {
Ok(self
.vars
.iter()
.find_map(|var| env::get_non_empty(var))
.map(|token| TransportAuth::Token {
username: Some(self.username.clone()),
token: SecretString::new(token),
}))
}
}
#[derive(Clone, Default)]
pub struct ChainAuthProvider {
providers: Vec<Arc<dyn AuthProvider>>,
}
impl ChainAuthProvider {
#[must_use]
pub fn new(providers: Vec<Arc<dyn AuthProvider>>) -> Self {
Self { providers }
}
#[must_use]
pub fn with(mut self, provider: Arc<dyn AuthProvider>) -> Self {
self.providers.push(provider);
self
}
}
impl std::fmt::Debug for ChainAuthProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ChainAuthProvider")
.field("providers", &self.providers.len())
.finish()
}
}
impl AuthProvider for ChainAuthProvider {
fn transport_auth(&self, remote: Option<&str>) -> AppResult<Option<TransportAuth>> {
for provider in &self.providers {
if let Some(auth) = provider.transport_auth(remote)? {
return Ok(Some(auth));
}
}
Ok(None)
}
fn signing_config(&self) -> AppResult<Option<SigningConfig>> {
for provider in &self.providers {
if let Some(config) = provider.signing_config()? {
return Ok(Some(config));
}
}
Ok(None)
}
}
#[cfg(test)]
mod tests {
use super::*;
const PRESENT_VAR: &str = "CARGO_PKG_NAME";
const PRESENT_VALUE: &str = "rskit-git";
const ABSENT_VAR: &str = "RSKIT_GIT_AUTH_TEST_ABSENT_VAR_9F3A";
#[test]
fn default_provider_offers_nothing() {
let provider = DefaultAuthProvider;
assert_eq!(provider.transport_auth(None).expect("resolve"), None);
assert!(provider.signing_config().expect("resolve").is_none());
}
#[test]
fn static_provider_returns_fixed_transport() {
let provider = StaticAuthProvider::new(TransportAuth::SshAgent {
username: "git".to_string(),
});
assert_eq!(
provider.transport_auth(Some("origin")).expect("resolve"),
Some(TransportAuth::SshAgent {
username: "git".to_string(),
})
);
}
#[test]
fn env_token_provider_reads_present_variable() {
let provider = EnvTokenAuthProvider::with_vars([ABSENT_VAR, PRESENT_VAR]);
let auth = provider.transport_auth(None).expect("resolve");
assert_eq!(
auth,
Some(TransportAuth::Token {
username: Some(DEFAULT_TOKEN_USERNAME.to_string()),
token: SecretString::new(PRESENT_VALUE),
})
);
}
#[test]
fn env_token_provider_honors_username_override() {
let provider = EnvTokenAuthProvider::with_var(PRESENT_VAR).with_username("token-user");
let auth = provider.transport_auth(None).expect("resolve");
assert_eq!(
auth,
Some(TransportAuth::Token {
username: Some("token-user".to_string()),
token: SecretString::new(PRESENT_VALUE),
})
);
}
#[test]
fn env_token_provider_absent_variable_is_none() {
let provider = EnvTokenAuthProvider::with_var(ABSENT_VAR);
assert_eq!(provider.transport_auth(None).expect("resolve"), None);
}
#[test]
fn chain_returns_first_some() {
let chain = ChainAuthProvider::new(vec![
Arc::new(EnvTokenAuthProvider::with_var(ABSENT_VAR)),
Arc::new(EnvTokenAuthProvider::with_var(PRESENT_VAR)),
Arc::new(StaticAuthProvider::new(TransportAuth::SshAgent {
username: "unused".to_string(),
})),
]);
let auth = chain.transport_auth(None).expect("resolve");
assert_eq!(
auth,
Some(TransportAuth::Token {
username: Some(DEFAULT_TOKEN_USERNAME.to_string()),
token: SecretString::new(PRESENT_VALUE),
})
);
}
#[test]
fn chain_falls_through_to_none() {
let chain = ChainAuthProvider::new(vec![
Arc::new(EnvTokenAuthProvider::with_var(ABSENT_VAR)),
Arc::new(DefaultAuthProvider),
]);
assert_eq!(chain.transport_auth(None).expect("resolve"), None);
assert!(chain.signing_config().expect("resolve").is_none());
}
}