1use oauth2::basic::{BasicClient, BasicRequestTokenError, BasicTokenResponse};
10use oauth2::{
11 AuthType, AuthUrl, ClientId, DeviceAuthorizationUrl, DeviceCodeErrorResponse,
12 DeviceCodeErrorResponseType, EndpointNotSet, EndpointSet, RefreshToken, RequestTokenError,
13 RevocationUrl, StandardRevocableToken, TokenResponse, TokenUrl,
14};
15use thiserror::Error;
16use url::Url;
17
18#[derive(Clone)]
22pub struct TokenSet {
23 pub access_token: String,
24 pub refresh_token: Option<String>,
25 pub expires_in: u64,
27}
28
29impl std::fmt::Debug for TokenSet {
30 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
31 f.debug_struct("TokenSet")
32 .field("access_token", &"<redacted>")
33 .field(
34 "refresh_token",
35 &self.refresh_token.as_ref().map(|_| "<redacted>"),
36 )
37 .field("expires_in", &self.expires_in)
38 .finish()
39 }
40}
41
42#[derive(Debug, Error)]
52#[non_exhaustive]
53pub enum AuthError {
54 #[error("the login code expired before it was approved; start login again")]
56 Expired,
57
58 #[error("the login request was denied")]
60 Denied,
61
62 #[error("network error contacting the identity provider: {0}")]
64 Network(#[from] reqwest::Error),
65
66 #[error("network error contacting the identity provider: {0}")]
71 Transport(String),
72
73 #[error("unexpected identity-provider response: {0}")]
75 Protocol(String),
76
77 #[error(
81 "this Redis Cloud account must be linked to social sign-in once before the CLI can use it"
82 )]
83 MigrationRequired,
84
85 #[error(
89 "your role on this Redis Cloud account cannot enable programmatic access; that needs \
90 {allowed_roles}. Ask someone who has it to enable it once in the console, then run \
91 login again"
92 )]
93 NotAccountOwner { allowed_roles: String },
94
95 #[error(
98 "programmatic access is not enabled for this Redis Cloud account; ask Redis support to \
99 enable API access for the account, then run login again"
100 )]
101 CapiDisabled,
102
103 #[error("account {requested} is not one of yours; you belong to: {available}")]
106 UnknownAccount { requested: u64, available: String },
107
108 #[error("{0}")]
111 AccountRequired(String),
112
113 #[error("this account requires multi-factor authentication")]
116 MfaRequired { factors: Vec<String> },
117
118 #[error("the multi-factor code was not accepted")]
120 MfaInvalidCode,
121
122 #[error("too many multi-factor attempts; wait before trying again")]
124 MfaQuotaExceeded,
125}
126
127pub(crate) type OktaClient =
130 BasicClient<EndpointSet, EndpointSet, EndpointNotSet, EndpointSet, EndpointSet>;
131
132pub(crate) fn endpoint(issuer: &Url, path: &str) -> String {
134 format!(
135 "{}/{}",
136 issuer.as_str().trim_end_matches('/'),
137 path.trim_start_matches('/')
138 )
139}
140
141pub(crate) fn okta_client(issuer: &Url, client_id: &str) -> Result<OktaClient, AuthError> {
148 let auth = AuthUrl::new(endpoint(issuer, "v1/authorize"))
149 .map_err(|e| AuthError::Protocol(format!("invalid authorize URL: {e}")))?;
150 let token = TokenUrl::new(endpoint(issuer, "v1/token"))
151 .map_err(|e| AuthError::Protocol(format!("invalid token URL: {e}")))?;
152 let device = DeviceAuthorizationUrl::new(endpoint(issuer, "v1/device/authorize"))
153 .map_err(|e| AuthError::Protocol(format!("invalid device-authorization URL: {e}")))?;
154 let revocation = RevocationUrl::new(endpoint(issuer, "v1/revoke"))
155 .map_err(|e| AuthError::Protocol(format!("invalid revocation URL: {e}")))?;
156 Ok(BasicClient::new(ClientId::new(client_id.to_string()))
157 .set_auth_uri(auth)
158 .set_token_uri(token)
159 .set_device_authorization_url(device)
160 .set_revocation_url(revocation)
161 .set_auth_type(AuthType::RequestBody))
162}
163
164pub(crate) fn oauth_http_client() -> Result<oauth2::reqwest::Client, AuthError> {
167 oauth2::reqwest::Client::builder()
168 .redirect(oauth2::reqwest::redirect::Policy::none())
169 .user_agent(crate::USER_AGENT)
170 .build()
171 .map_err(|e| AuthError::Protocol(format!("could not build the OAuth HTTP client: {e}")))
172}
173
174pub(crate) fn default_http_client() -> reqwest::Client {
176 reqwest::Client::builder()
177 .user_agent(crate::USER_AGENT)
178 .redirect(reqwest::redirect::Policy::none())
179 .build()
180 .expect("building the reqwest client should not fail")
181}
182
183pub(crate) async fn revoke_refresh_token(
188 issuer: &Url,
189 client_id: &str,
190 refresh_token: &str,
191) -> Result<(), AuthError> {
192 let client = okta_client(issuer, client_id)?;
193 let http = oauth_http_client()?;
194 client
195 .revoke_token(StandardRevocableToken::RefreshToken(RefreshToken::new(
196 refresh_token.to_string(),
197 )))
198 .map_err(|e| AuthError::Protocol(format!("could not build the revocation request: {e}")))?
199 .request_async(&http)
200 .await
201 .map_err(|e| AuthError::Protocol(format!("token revocation failed: {e}")))?;
202 Ok(())
203}
204
205pub(crate) fn to_token_set(resp: &BasicTokenResponse) -> TokenSet {
207 TokenSet {
208 access_token: resp.access_token().secret().clone(),
209 refresh_token: resp.refresh_token().map(|r| r.secret().clone()),
210 expires_in: resp.expires_in().map(|d| d.as_secs()).unwrap_or(0),
211 }
212}
213
214pub(crate) fn map_basic_token_error<RE>(err: BasicRequestTokenError<RE>) -> AuthError
220where
221 RE: std::error::Error,
222{
223 match err {
224 RequestTokenError::ServerResponse(resp) => match resp.error().as_ref() {
225 "access_denied" => AuthError::Denied,
226 "expired_token" => AuthError::Expired,
227 _ => AuthError::Protocol(format!("identity-provider error: {resp}")),
228 },
229 RequestTokenError::Request(e) => AuthError::Transport(error_chain(&e)),
230 other => AuthError::Protocol(other.to_string()),
231 }
232}
233
234pub(crate) fn map_device_token_error<RE>(
238 err: RequestTokenError<RE, DeviceCodeErrorResponse>,
239) -> AuthError
240where
241 RE: std::error::Error,
242{
243 match err {
244 RequestTokenError::ServerResponse(resp) => match resp.error() {
245 DeviceCodeErrorResponseType::ExpiredToken => AuthError::Expired,
246 DeviceCodeErrorResponseType::AccessDenied => AuthError::Denied,
247 _ => AuthError::Protocol(format!("identity-provider error: {resp}")),
248 },
249 RequestTokenError::Request(e) => AuthError::Transport(error_chain(&e)),
250 other => AuthError::Protocol(other.to_string()),
251 }
252}
253
254fn error_chain(err: &dyn std::error::Error) -> String {
259 let mut parts = vec![err.to_string()];
260 let mut source = err.source();
261 while let Some(e) = source {
262 let text = e.to_string();
263 if !parts.iter().any(|p| p == &text) {
264 parts.push(text);
265 }
266 source = e.source();
267 }
268 parts.join(": ")
269}
270
271pub(crate) fn truncate(s: &str) -> String {
273 const MAX: usize = 200;
274 if s.chars().count() <= MAX {
275 s.to_string()
276 } else {
277 let head: String = s.chars().take(MAX).collect();
278 format!("{head}…")
279 }
280}
281
282pub(crate) async fn refresh(
288 issuer: &Url,
289 client_id: &str,
290 refresh_token: &str,
291) -> Result<TokenSet, AuthError> {
292 let client = okta_client(issuer, client_id)?;
293 let http = oauth_http_client()?;
294 let resp = client
295 .exchange_refresh_token(&RefreshToken::new(refresh_token.to_string()))
296 .request_async(&http)
297 .await
298 .map_err(map_basic_token_error)?;
299 Ok(to_token_set(&resp))
300}
301
302#[cfg(test)]
303mod tests {
304 use super::*;
305 use wiremock::matchers::{method, path};
306 use wiremock::{Mock, MockServer, ResponseTemplate};
307
308 async fn mount_token(server: &MockServer, status: u16, body: serde_json::Value) {
309 Mock::given(method("POST"))
310 .and(path("/v1/token"))
311 .respond_with(ResponseTemplate::new(status).set_body_json(body))
312 .mount(server)
313 .await;
314 }
315
316 #[tokio::test]
317 async fn refresh_returns_rotated_token() {
318 let server = MockServer::start().await;
319 mount_token(
320 &server,
321 200,
322 serde_json::json!({
323 "access_token": "AT2",
324 "token_type": "Bearer",
325 "refresh_token": "RT2",
326 "expires_in": 3600
327 }),
328 )
329 .await;
330
331 let issuer = Url::parse(&server.uri()).unwrap();
332 let t = refresh(&issuer, "test-client", "RT1").await.unwrap();
333 assert_eq!(t.access_token, "AT2");
334 assert_eq!(t.refresh_token.as_deref(), Some("RT2"));
336 assert_eq!(t.expires_in, 3600);
337 }
338
339 #[tokio::test]
340 async fn refresh_error_is_protocol() {
341 let server = MockServer::start().await;
342 mount_token(
343 &server,
344 400,
345 serde_json::json!({"error": "invalid_grant", "error_description": "expired"}),
346 )
347 .await;
348 let issuer = Url::parse(&server.uri()).unwrap();
349 assert!(matches!(
350 refresh(&issuer, "test-client", "RT1").await,
351 Err(AuthError::Protocol(_))
352 ));
353 }
354
355 #[tokio::test]
359 async fn refresh_transport_failure_is_transport_not_protocol() {
360 let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
361 let port = listener.local_addr().unwrap().port();
362 drop(listener);
363
364 let issuer = Url::parse(&format!("http://127.0.0.1:{port}")).unwrap();
365 let err = refresh(&issuer, "test-client", "RT1").await.unwrap_err();
366 assert!(matches!(err, AuthError::Transport(_)), "got {err:?}");
367 }
368
369 #[test]
370 fn token_set_debug_redacts_secrets() {
371 let t = TokenSet {
372 access_token: "AT-should-not-appear".into(),
373 refresh_token: Some("RT-should-not-appear".into()),
374 expires_in: 3600,
375 };
376 let dbg = format!("{t:?}");
377 assert!(dbg.contains("<redacted>"));
378 assert!(!dbg.contains("AT-should-not-appear"));
379 assert!(!dbg.contains("RT-should-not-appear"));
380 assert!(dbg.contains("3600"));
381 }
382}