use oauth2::basic::{BasicClient, BasicRequestTokenError, BasicTokenResponse};
use oauth2::{
AuthType, AuthUrl, ClientId, DeviceAuthorizationUrl, DeviceCodeErrorResponse,
DeviceCodeErrorResponseType, EndpointNotSet, EndpointSet, RefreshToken, RequestTokenError,
RevocationUrl, StandardRevocableToken, TokenResponse, TokenUrl,
};
use thiserror::Error;
use url::Url;
#[derive(Clone)]
pub struct TokenSet {
pub access_token: String,
pub refresh_token: Option<String>,
pub expires_in: u64,
}
impl std::fmt::Debug for TokenSet {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TokenSet")
.field("access_token", &"<redacted>")
.field(
"refresh_token",
&self.refresh_token.as_ref().map(|_| "<redacted>"),
)
.field("expires_in", &self.expires_in)
.finish()
}
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum AuthError {
#[error("the login code expired before it was approved; start login again")]
Expired,
#[error("the login request was denied")]
Denied,
#[error("network error contacting the identity provider: {0}")]
Network(#[from] reqwest::Error),
#[error("network error contacting the identity provider: {0}")]
Transport(String),
#[error("unexpected identity-provider response: {0}")]
Protocol(String),
#[error(
"this Redis Cloud account must be linked to social sign-in once before the CLI can use it"
)]
MigrationRequired,
#[error(
"your role on this Redis Cloud account cannot enable programmatic access; that needs \
{allowed_roles}. Ask someone who has it to enable it once in the console, then run \
login again"
)]
NotAccountOwner { allowed_roles: String },
#[error(
"programmatic access is not enabled for this Redis Cloud account; ask Redis support to \
enable API access for the account, then run login again"
)]
CapiDisabled,
#[error("account {requested} is not one of yours; you belong to: {available}")]
UnknownAccount { requested: u64, available: String },
#[error("{0}")]
AccountRequired(String),
#[error("this account requires multi-factor authentication")]
MfaRequired { factors: Vec<String> },
#[error("the multi-factor code was not accepted")]
MfaInvalidCode,
#[error("too many multi-factor attempts; wait before trying again")]
MfaQuotaExceeded,
}
pub(crate) type OktaClient =
BasicClient<EndpointSet, EndpointSet, EndpointNotSet, EndpointSet, EndpointSet>;
pub(crate) fn endpoint(issuer: &Url, path: &str) -> String {
format!(
"{}/{}",
issuer.as_str().trim_end_matches('/'),
path.trim_start_matches('/')
)
}
pub(crate) fn okta_client(issuer: &Url, client_id: &str) -> Result<OktaClient, AuthError> {
let auth = AuthUrl::new(endpoint(issuer, "v1/authorize"))
.map_err(|e| AuthError::Protocol(format!("invalid authorize URL: {e}")))?;
let token = TokenUrl::new(endpoint(issuer, "v1/token"))
.map_err(|e| AuthError::Protocol(format!("invalid token URL: {e}")))?;
let device = DeviceAuthorizationUrl::new(endpoint(issuer, "v1/device/authorize"))
.map_err(|e| AuthError::Protocol(format!("invalid device-authorization URL: {e}")))?;
let revocation = RevocationUrl::new(endpoint(issuer, "v1/revoke"))
.map_err(|e| AuthError::Protocol(format!("invalid revocation URL: {e}")))?;
Ok(BasicClient::new(ClientId::new(client_id.to_string()))
.set_auth_uri(auth)
.set_token_uri(token)
.set_device_authorization_url(device)
.set_revocation_url(revocation)
.set_auth_type(AuthType::RequestBody))
}
pub(crate) fn oauth_http_client() -> Result<oauth2::reqwest::Client, AuthError> {
oauth2::reqwest::Client::builder()
.redirect(oauth2::reqwest::redirect::Policy::none())
.user_agent(crate::USER_AGENT)
.build()
.map_err(|e| AuthError::Protocol(format!("could not build the OAuth HTTP client: {e}")))
}
pub(crate) fn default_http_client() -> reqwest::Client {
reqwest::Client::builder()
.user_agent(crate::USER_AGENT)
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("building the reqwest client should not fail")
}
pub(crate) async fn revoke_refresh_token(
issuer: &Url,
client_id: &str,
refresh_token: &str,
) -> Result<(), AuthError> {
let client = okta_client(issuer, client_id)?;
let http = oauth_http_client()?;
client
.revoke_token(StandardRevocableToken::RefreshToken(RefreshToken::new(
refresh_token.to_string(),
)))
.map_err(|e| AuthError::Protocol(format!("could not build the revocation request: {e}")))?
.request_async(&http)
.await
.map_err(|e| AuthError::Protocol(format!("token revocation failed: {e}")))?;
Ok(())
}
pub(crate) fn to_token_set(resp: &BasicTokenResponse) -> TokenSet {
TokenSet {
access_token: resp.access_token().secret().clone(),
refresh_token: resp.refresh_token().map(|r| r.secret().clone()),
expires_in: resp.expires_in().map(|d| d.as_secs()).unwrap_or(0),
}
}
pub(crate) fn map_basic_token_error<RE>(err: BasicRequestTokenError<RE>) -> AuthError
where
RE: std::error::Error,
{
match err {
RequestTokenError::ServerResponse(resp) => match resp.error().as_ref() {
"access_denied" => AuthError::Denied,
"expired_token" => AuthError::Expired,
_ => AuthError::Protocol(format!("identity-provider error: {resp}")),
},
RequestTokenError::Request(e) => AuthError::Transport(error_chain(&e)),
other => AuthError::Protocol(other.to_string()),
}
}
pub(crate) fn map_device_token_error<RE>(
err: RequestTokenError<RE, DeviceCodeErrorResponse>,
) -> AuthError
where
RE: std::error::Error,
{
match err {
RequestTokenError::ServerResponse(resp) => match resp.error() {
DeviceCodeErrorResponseType::ExpiredToken => AuthError::Expired,
DeviceCodeErrorResponseType::AccessDenied => AuthError::Denied,
_ => AuthError::Protocol(format!("identity-provider error: {resp}")),
},
RequestTokenError::Request(e) => AuthError::Transport(error_chain(&e)),
other => AuthError::Protocol(other.to_string()),
}
}
fn error_chain(err: &dyn std::error::Error) -> String {
let mut parts = vec![err.to_string()];
let mut source = err.source();
while let Some(e) = source {
let text = e.to_string();
if !parts.iter().any(|p| p == &text) {
parts.push(text);
}
source = e.source();
}
parts.join(": ")
}
pub(crate) fn truncate(s: &str) -> String {
const MAX: usize = 200;
if s.chars().count() <= MAX {
s.to_string()
} else {
let head: String = s.chars().take(MAX).collect();
format!("{head}…")
}
}
pub(crate) async fn refresh(
issuer: &Url,
client_id: &str,
refresh_token: &str,
) -> Result<TokenSet, AuthError> {
let client = okta_client(issuer, client_id)?;
let http = oauth_http_client()?;
let resp = client
.exchange_refresh_token(&RefreshToken::new(refresh_token.to_string()))
.request_async(&http)
.await
.map_err(map_basic_token_error)?;
Ok(to_token_set(&resp))
}
#[cfg(test)]
mod tests {
use super::*;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
async fn mount_token(server: &MockServer, status: u16, body: serde_json::Value) {
Mock::given(method("POST"))
.and(path("/v1/token"))
.respond_with(ResponseTemplate::new(status).set_body_json(body))
.mount(server)
.await;
}
#[tokio::test]
async fn refresh_returns_rotated_token() {
let server = MockServer::start().await;
mount_token(
&server,
200,
serde_json::json!({
"access_token": "AT2",
"token_type": "Bearer",
"refresh_token": "RT2",
"expires_in": 3600
}),
)
.await;
let issuer = Url::parse(&server.uri()).unwrap();
let t = refresh(&issuer, "test-client", "RT1").await.unwrap();
assert_eq!(t.access_token, "AT2");
assert_eq!(t.refresh_token.as_deref(), Some("RT2"));
assert_eq!(t.expires_in, 3600);
}
#[tokio::test]
async fn refresh_error_is_protocol() {
let server = MockServer::start().await;
mount_token(
&server,
400,
serde_json::json!({"error": "invalid_grant", "error_description": "expired"}),
)
.await;
let issuer = Url::parse(&server.uri()).unwrap();
assert!(matches!(
refresh(&issuer, "test-client", "RT1").await,
Err(AuthError::Protocol(_))
));
}
#[tokio::test]
async fn refresh_transport_failure_is_transport_not_protocol() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
drop(listener);
let issuer = Url::parse(&format!("http://127.0.0.1:{port}")).unwrap();
let err = refresh(&issuer, "test-client", "RT1").await.unwrap_err();
assert!(matches!(err, AuthError::Transport(_)), "got {err:?}");
}
#[test]
fn token_set_debug_redacts_secrets() {
let t = TokenSet {
access_token: "AT-should-not-appear".into(),
refresh_token: Some("RT-should-not-appear".into()),
expires_in: 3600,
};
let dbg = format!("{t:?}");
assert!(dbg.contains("<redacted>"));
assert!(!dbg.contains("AT-should-not-appear"));
assert!(!dbg.contains("RT-should-not-appear"));
assert!(dbg.contains("3600"));
}
}