use std::sync::Arc;
use anyhow::{Context, Result};
use aviso::ClientError;
use std::path::Path;
use aviso::auth::{AuthProvider, Basic, Bearer};
pub(crate) fn provider_from_flags(
token: Option<&str>,
username: Option<&str>,
password: Option<&str>,
) -> Result<Option<Arc<dyn AuthProvider>>> {
if let Some(t) = token {
let bearer =
Bearer::new(t.to_string()).context("build Bearer auth provider from --token flag")?;
return Ok(Some(Arc::new(bearer)));
}
if let (Some(u), Some(p)) = (username, password) {
let basic = Basic::new(u.to_string(), p.to_string())
.context("build Basic auth provider from --username/--password flags")?;
return Ok(Some(Arc::new(basic)));
}
Ok(None)
}
pub(crate) type SelectedProvider = (Option<Arc<dyn AuthProvider>>, Option<&'static str>);
pub(crate) fn resolve_provider(
flag_provider: Option<Arc<dyn AuthProvider>>,
config_path: &Path,
config_content: Option<String>,
) -> Result<SelectedProvider> {
if let Some(provider) = flag_provider {
return Ok((Some(provider), Some("flag")));
}
let mut paths = aviso::auth::DiscoveryPaths::from_env();
match config_content {
Some(content) => {
paths.config_file = Some(config_path.to_path_buf());
paths.config_content = Some(content);
}
None => paths.config_file = None,
}
let found = aviso::auth::discover_with(&paths).map_err(|e| match e {
ClientError::Auth(reason) => crate::exit::usage_error(reason),
other => anyhow::Error::from(other),
})?;
Ok(match found {
Some(found) => {
let label = found.source().label();
(Some(found.into_provider()), Some(label))
}
None => (None, None),
})
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
reason = "test code: unwrap/expect on provider construction is the expected diagnostic"
)]
mod tests {
use super::*;
#[test]
fn an_absent_snapshot_does_not_reopen_a_config_file_that_appeared_later() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("config.yaml");
std::fs::write(&path, "auth:\n bearer_token: appeared-later\n").expect("write");
let _isolate = IsolatedSources::new(dir.path());
let (provider, source) = resolve_provider(None, &path, None).expect("resolve");
assert!(provider.is_none(), "the later file must not be read");
assert!(source.is_none());
}
#[test]
fn a_present_snapshot_is_used_even_if_the_file_changed_since() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("config.yaml");
std::fs::write(&path, "auth:\n bearer_token: on-disk-now\n").expect("write");
let _isolate = IsolatedSources::new(dir.path());
let snapshot = "auth:\n bearer_token: from-snapshot\n".to_string();
let (provider, source) = resolve_provider(None, &path, Some(snapshot)).expect("resolve");
assert!(provider.is_some());
assert_eq!(source, Some("config file"));
}
struct IsolatedSources {
_guard: std::sync::MutexGuard<'static, ()>,
saved: Vec<(&'static str, Option<std::ffi::OsString>)>,
}
static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
impl IsolatedSources {
fn new(dir: &Path) -> Self {
let guard = ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let names = [
"AVISO_TOKEN",
"AVISO_USERNAME",
"AVISO_PASSWORD",
"AVISO_CREDENTIALS_FILE",
];
let saved = names.iter().map(|k| (*k, std::env::var_os(k))).collect();
unsafe {
for name in &names[..3] {
std::env::remove_var(name);
}
std::env::set_var(
"AVISO_CREDENTIALS_FILE",
dir.join("absent-credentials.yaml"),
);
}
Self {
_guard: guard,
saved,
}
}
}
impl Drop for IsolatedSources {
fn drop(&mut self) {
unsafe {
for (name, value) in &self.saved {
match value {
Some(v) => std::env::set_var(name, v),
None => std::env::remove_var(name),
}
}
}
}
}
#[test]
fn flag_token_yields_bearer_provider() {
let p = provider_from_flags(Some("the-token"), None, None).unwrap();
assert!(p.is_some());
}
#[test]
fn flag_username_password_yield_basic_provider() {
let p = provider_from_flags(None, Some("alice"), Some("hunter2")).unwrap();
assert!(p.is_some());
}
#[test]
fn flag_username_only_yields_none() {
let p = provider_from_flags(None, Some("alice"), None).unwrap();
assert!(p.is_none());
}
#[test]
fn flag_empty_token_errors() {
let err = provider_from_flags(Some(""), None, None).unwrap_err();
let s = err.to_string();
assert!(
s.contains("Bearer") || s.contains("--token"),
"error should name the source: {s}"
);
}
#[test]
fn empty_flag_yields_none() {
let p = provider_from_flags(None, None, None).unwrap();
assert!(p.is_none());
}
}