use std::collections::HashMap;
use std::sync::Arc;
use std::sync::Mutex as StdMutex;
use std::sync::Weak;
use tokio::sync::Mutex as AsyncMutex;
use crate::Error;
use crate::Res;
use crate::error::AuthError;
use crate::error::LoginError;
use crate::io::remote::client::HttpClient;
use crate::io::storage::LocalStorage;
use crate::io::storage::Storage;
use crate::io::storage::auth::AuthIo;
use crate::io::storage::auth::Credentials;
use crate::io::storage::auth::OAuthClient;
use crate::io::storage::auth::Tokens;
use crate::paths::DomainPaths;
use quilt_uri::Host;
use tracing::debug;
use tracing::error;
use tracing::info;
use tracing::warn;
mod graphql;
mod oauth;
mod registry;
mod retry;
pub use oauth::OAuthParams;
pub use oauth::PkceChallenge;
pub use oauth::catalog_authorize_url;
pub use oauth::connect_host;
pub use oauth::pkce_challenge;
pub use oauth::random_state;
pub use registry::RemoteTokens;
use graphql::mutate_switch_role;
use graphql::query_buckets;
use graphql::query_me;
use oauth::exchange_oauth_code;
use oauth::refresh_oauth_tokens;
use oauth::register_client;
use registry::get_auth_tokens;
use registry::get_registry_url;
use registry::refresh_credentials;
use retry::classify_retry_outcome;
use retry::http_status;
use retry::is_credentials_auth_error;
use retry::is_role_auth_error;
use retry::is_token_auth_error;
#[cfg(test)]
mod test_utils;
#[cfg(test)]
mod tests;
type RefreshLocks = Arc<StdMutex<HashMap<Host, Weak<AsyncMutex<()>>>>>;
const ROLE_ENDPOINT: &str = "registry GraphQL endpoint";
type SessionRoles = Arc<StdMutex<HashMap<Host, String>>>;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RoleInfo {
pub current: String,
pub available: Vec<String>,
}
impl From<graphql::Me> for RoleInfo {
fn from(me: graphql::Me) -> Self {
Self {
current: me.role.name,
available: me.roles.into_iter().map(|r| r.name).collect(),
}
}
}
#[derive(Debug)]
pub struct Auth<S: Storage = LocalStorage> {
pub paths: DomainPaths,
pub storage: Arc<S>,
refresh_locks: RefreshLocks,
session_roles: SessionRoles,
}
impl<S: Storage> Clone for Auth<S> {
fn clone(&self) -> Self {
Self {
paths: self.paths.clone(),
storage: Arc::clone(&self.storage),
refresh_locks: Arc::clone(&self.refresh_locks),
session_roles: Arc::clone(&self.session_roles),
}
}
}
impl<S: Storage + Send + Sync> Auth<S> {
pub fn new(paths: DomainPaths, storage: Arc<S>) -> Self {
Self {
paths,
storage,
refresh_locks: Arc::new(StdMutex::new(HashMap::new())),
session_roles: Arc::new(StdMutex::new(HashMap::new())),
}
}
fn refresh_lock_for(&self, host: &Host) -> Arc<AsyncMutex<()>> {
let mut locks = self
.refresh_locks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
locks.retain(|_, weak| weak.strong_count() > 0);
if let Some(arc) = locks.get(host).and_then(Weak::upgrade) {
return arc;
}
let arc = Arc::new(AsyncMutex::new(()));
locks.insert(host.clone(), Arc::downgrade(&arc));
arc
}
pub async fn login<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
refresh_token: String,
) -> Res {
info!("⏳ Logging in to host {} with refresh token", host);
let tokens = match self
.get_auth_tokens(http_client, host, &refresh_token)
.await
{
Ok(t) => t,
Err(e) => {
warn!("❌ Failed to get auth tokens for {}: {}", host, e);
return Err(e);
}
};
if let Err(e) = self.save_tokens(host, &tokens).await {
warn!("❌ Failed to save tokens for {}: {}", host, e);
return Err(e);
}
if let Err(e) = self
.refresh_credentials(http_client, host, &tokens.access_token)
.await
{
warn!("❌ Failed to refresh credentials for {}: {}", host, e);
return Err(e);
}
info!("✔️ Successfully logged in and authenticated to {}", host);
Ok(())
}
pub async fn get_or_register_client<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
redirect_uri: &str,
) -> Res<OAuthClient> {
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
if let Some(client) = auth_io.read_client().await? {
if client.redirect_uri == redirect_uri {
info!("✔️ Found existing OAuth client for {}", host);
return Ok(client);
}
info!(
"⚠️ Cached client has stale redirect_uri, re-registering for {}",
host
);
}
info!("⏳ Registering new OAuth client for {}", host);
let client = register_client(http_client, host, redirect_uri).await?;
auth_io.write_client(&client).await?;
info!(
"✔️ Registered OAuth client for {}: {}",
host, client.client_id
);
Ok(client)
}
pub async fn login_oauth<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
params: OAuthParams,
) -> Res {
info!("⏳ OAuth login for host {}", host);
let tokens = exchange_oauth_code(http_client, host, ¶ms)
.await
.map_err(|e| {
warn!("❌ Failed to exchange OAuth code for {}: {}", host, e);
e
})?;
self.save_tokens(host, &tokens).await.map_err(|e| {
warn!("❌ Failed to save tokens for {}: {}", host, e);
e
})?;
self.refresh_credentials(http_client, host, &tokens.access_token)
.await
.map_err(|e| {
warn!("❌ Failed to refresh credentials for {}: {}", host, e);
e
})?;
info!("✔️ OAuth login successful for {}", host);
Ok(())
}
async fn get_auth_tokens<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
refresh_token: &str,
) -> Res<Tokens> {
debug!("⏳ Getting auth tokens for host {:?}", host);
let tokens = get_auth_tokens(http_client, host, refresh_token).await?;
debug!("✔️ Successfully retrieved auth tokens");
Ok(tokens)
}
async fn save_tokens(&self, host: &Host, tokens: &Tokens) -> Res<()> {
debug!("⏳ Saving tokens for host {:?}", host);
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
auth_io.write_tokens(tokens).await?;
debug!(
"✔️ Successfully saved tokens to the {:?}",
self.paths.auth_host(host)
);
Ok(())
}
async fn refresh_tokens<T: HttpClient>(
&self,
http_client: &T,
auth_io: &AuthIo<Arc<S>>,
host: &Host,
tokens: &Tokens,
) -> Res<Tokens> {
let client = auth_io
.read_client()
.await?
.ok_or(LoginError::Required(Some(host.to_owned())))?;
let new_tokens =
refresh_oauth_tokens(http_client, host, &tokens.refresh_token, &client.client_id)
.await?;
auth_io.write_tokens(&new_tokens).await?;
info!("✔️ Successfully refreshed tokens for {}", host);
Ok(new_tokens)
}
async fn refresh_tokens_with_retry<T: HttpClient>(
&self,
http_client: &T,
auth_io: &AuthIo<Arc<S>>,
host: &Host,
tokens: &Tokens,
) -> Res<Tokens> {
let first_err = match self
.refresh_tokens(http_client, auth_io, host, tokens)
.await
{
Ok(t) => return Ok(t),
Err(e) => e,
};
if matches!(first_err, Error::Login(LoginError::Required(_))) {
warn!("❌ No OAuth client registered for {}, login required", host);
return Err(first_err);
}
if !is_token_auth_error(&first_err) {
warn!(
status = ?http_status(&first_err),
"❌ Failed to refresh tokens for {}: {}", host, first_err
);
return Err(first_err);
}
info!(
status = ?http_status(&first_err),
"⚠️ Auth error refreshing tokens for {}, retrying once: {}", host, first_err
);
classify_retry_outcome(
self.refresh_tokens(http_client, auth_io, host, tokens)
.await,
is_token_auth_error,
"token endpoint",
host,
)
}
async fn refresh_credentials_with_retry<T: HttpClient>(
&self,
http_client: &T,
auth_io: &AuthIo<Arc<S>>,
host: &Host,
access_token: &str,
) -> Res<Credentials> {
let first_err = match self
.refresh_credentials(http_client, host, access_token)
.await
{
Ok(c) => return Ok(c),
Err(e) => e,
};
if !is_credentials_auth_error(&first_err) {
warn!(
status = ?http_status(&first_err),
"❌ Failed to refresh credentials for {}: {}", host, first_err
);
return Err(first_err);
}
info!(
status = ?http_status(&first_err),
"⚠️ Auth error refreshing credentials for {}, \
force-refreshing token and retrying: {}",
host, first_err
);
let access_token = self
.force_refresh_access_token(http_client, auth_io, host)
.await?;
classify_retry_outcome(
self.refresh_credentials(http_client, host, &access_token)
.await,
is_credentials_auth_error,
"credentials endpoint",
host,
)
}
async fn force_refresh_access_token<T: HttpClient>(
&self,
http_client: &T,
auth_io: &AuthIo<Arc<S>>,
host: &Host,
) -> Res<String> {
let tokens = auth_io
.read_tokens()
.await?
.ok_or_else(|| LoginError::Required(Some(host.to_owned())))?;
let new_tokens = self
.refresh_tokens_with_retry(http_client, auth_io, host, &tokens)
.await?;
Ok(new_tokens.access_token)
}
async fn refresh_credentials<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
access_token: &str,
) -> Res<Credentials> {
debug!("⏳ Refreshing credentials for host {:?}", host);
let credentials = refresh_credentials(http_client, host, access_token).await?;
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
auth_io.write_credentials(&credentials).await?;
debug!(
"✔️ Successfully refreshed credentials in {:?}",
self.paths.auth_host(host)
);
Ok(credentials)
}
pub async fn expire_credentials(&self, host: &Host) -> Res {
let lock = self.refresh_lock_for(host);
let _guard = lock.lock().await;
self.expire_credentials_locked(host).await
}
async fn expire_credentials_locked(&self, host: &Host) -> Res {
info!("⏳ Expiring cached credentials for {}", host);
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
auth_io.delete_credentials().await?;
info!("✔️ Cached credentials expired for {}", host);
Ok(())
}
fn observe_role(&self, host: &Host, role: &str) -> bool {
let mut roles = self
.session_roles
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
!matches!(roles.insert(host.clone(), role.to_owned()), Some(prev) if prev == role)
}
fn forget_role(&self, host: &Host) {
self.session_roles
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(host);
}
async fn reconcile_role(&self, host: &Host, role: &str) -> Res {
if !self.observe_role(host, role) {
return Ok(());
}
info!(
"⚠️ Active role for {} is {}, flushing credentials",
host, role
);
self.expire_credentials_locked(host).await.inspect_err(|_| {
warn!(
"❌ Flush for {} failed; keeping the role unknown so the next \
observation retries",
host
);
self.forget_role(host);
})
}
async fn valid_access_token<T: HttpClient>(&self, http_client: &T, host: &Host) -> Res<String> {
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
let tokens = auth_io
.read_tokens()
.await?
.ok_or_else(|| LoginError::Required(Some(host.to_owned())))?;
if tokens.expires_at <= chrono::Utc::now() + chrono::Duration::seconds(60) {
info!("⏳ Access token expired for {}, refreshing", host);
let refreshed = self
.refresh_tokens_with_retry(http_client, &auth_io, host, &tokens)
.await?;
return Ok(refreshed.access_token);
}
Ok(tokens.access_token)
}
async fn role_call_context<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
) -> Res<(url::Host, String)> {
let access_token = self.valid_access_token(http_client, host).await?;
let registry = get_registry_url(http_client, host).await?;
Ok((registry, access_token))
}
async fn role_retry_token<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
first_err: Error,
) -> Res<String> {
if !is_role_auth_error(&first_err) {
return Err(first_err);
}
info!(
status = ?http_status(&first_err),
"⚠️ Auth error on the role surface for {}, \
force-refreshing token and retrying: {}",
host, first_err
);
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
self.force_refresh_access_token(http_client, &auth_io, host)
.await
}
pub async fn refresh_roles<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
) -> Res<RoleInfo> {
let lock = self.refresh_lock_for(host);
let _guard = lock.lock().await;
let (registry, access_token) = self.role_call_context(http_client, host).await?;
let me = match query_me(http_client, ®istry, host, &access_token).await {
Ok(me) => me,
Err(first_err) => {
let retry_token = self.role_retry_token(http_client, host, first_err).await?;
classify_retry_outcome(
query_me(http_client, ®istry, host, &retry_token).await,
is_role_auth_error,
ROLE_ENDPOINT,
host,
)?
}
};
debug!("🔎 Asked {} for the active role: {}", host, me.role.name);
self.reconcile_role(host, &me.role.name).await?;
Ok(RoleInfo::from(me))
}
pub async fn switch_role<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
role_name: &str,
) -> Res<RoleInfo> {
info!("⏳ Switching role on {} to {}", host, role_name);
let lock = self.refresh_lock_for(host);
let _guard = lock.lock().await;
let (registry, access_token) = self.role_call_context(http_client, host).await?;
let me = match mutate_switch_role(http_client, ®istry, role_name, &access_token).await {
Ok(me) => me,
Err(first_err) => {
let retry_token = self.role_retry_token(http_client, host, first_err).await?;
classify_retry_outcome(
mutate_switch_role(http_client, ®istry, role_name, &retry_token).await,
is_role_auth_error,
ROLE_ENDPOINT,
host,
)?
}
};
self.forget_role(host);
self.reconcile_role(host, &me.role.name).await?;
info!("✔️ Switched role on {} to {}", host, me.role.name);
Ok(RoleInfo::from(me))
}
pub async fn readable_buckets<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
) -> Res<Vec<String>> {
let lock = self.refresh_lock_for(host);
let _guard = lock.lock().await;
let (registry, access_token) = self.role_call_context(http_client, host).await?;
match query_buckets(http_client, ®istry, &access_token).await {
Ok(buckets) => Ok(buckets),
Err(first_err) => {
let retry_token = self.role_retry_token(http_client, host, first_err).await?;
classify_retry_outcome(
query_buckets(http_client, ®istry, &retry_token).await,
is_role_auth_error,
ROLE_ENDPOINT,
host,
)
}
}
}
pub async fn get_credentials_or_refresh<T: HttpClient>(
&self,
http_client: &T,
host: &Host,
) -> Res<Credentials> {
info!("⏳ Getting or refreshing credentials for {}", host);
let auth_io = AuthIo::new(self.storage.clone(), self.paths.auth_host(host));
match auth_io.read_credentials().await {
Ok(Some(creds)) => {
debug!("✔️ Found valid credentials for {}", host);
return Ok(creds);
}
Ok(None) => {
info!("❌ No existing credentials found for {}", host);
}
Err(e) => {
error!("❌ Failed to read credentials for {}: {}", host, e);
return Err(Error::Auth(
host.to_owned(),
AuthError::CredentialsRead(e.to_string()),
));
}
}
let lock = self.refresh_lock_for(host);
let _guard = lock.lock().await;
match auth_io.read_credentials().await {
Ok(Some(creds)) => {
debug!("✔️ Another task refreshed credentials for {}", host);
return Ok(creds);
}
Ok(None) => {}
Err(e) => {
error!("❌ Failed to re-read credentials for {}: {}", host, e);
return Err(Error::Auth(
host.to_owned(),
AuthError::CredentialsRead(e.to_string()),
));
}
}
match auth_io.read_tokens().await {
Ok(Some(_)) => {}
Ok(None) => {
warn!("❌ No tokens found for {}, login required", host);
return Err(LoginError::Required(Some(host.to_owned())).into());
}
Err(e) => {
error!("❌ Failed to read tokens for {}: {}", host, e);
return Err(Error::Auth(
host.to_owned(),
AuthError::TokensRead(e.to_string()),
));
}
}
let access_token = self.valid_access_token(http_client, host).await?;
info!("⏳ Refreshing credentials using access token for {}", host);
let creds = self
.refresh_credentials_with_retry(http_client, &auth_io, host, &access_token)
.await?;
info!("✔️ Successfully refreshed credentials for {}", host);
Ok(creds)
}
}