1use std::time::Duration;
16
17use base64::Engine;
18use serde::Deserialize;
19use tokio::time::sleep;
20
21use crate::error::{CliError, Result};
22use crate::util;
23
24const MAX_TOKEN_RETRIES: u32 = 3;
28
29#[derive(Debug, Clone, Deserialize)]
31pub struct OAuth2Error {
32 pub error: String,
33 pub error_description: Option<String>,
34 pub error_uri: Option<String>,
35}
36
37pub fn redact_token_fields(s: &str) -> String {
44 const TOKEN_FIELDS: &[&str] = &["access_token", "refresh_token", "id_token"];
45 let mut result = s.to_string();
46 for field in TOKEN_FIELDS {
47 let key_pattern = format!("\"{field}\"");
49 let mut search_from = 0;
50 while let Some(key_pos) = result[search_from..].find(&key_pattern) {
51 let abs_key_end = search_from + key_pos + key_pattern.len();
52 let after_key = result[abs_key_end..].trim_start();
54 if !after_key.starts_with(':') {
55 search_from = abs_key_end;
56 continue;
57 }
58 let colon_offset = result[abs_key_end..].len() - after_key.len() + 1;
59 let abs_colon_end = abs_key_end + colon_offset;
60 let after_colon = result[abs_colon_end..].trim_start();
61 if !after_colon.starts_with('"') {
62 search_from = abs_colon_end;
63 continue;
64 }
65 let value_start_offset = result[abs_colon_end..].len() - after_colon.len();
67 let abs_value_open = abs_colon_end + value_start_offset; let value_chars = &result[abs_value_open + 1..]; let mut value_len = 0;
70 let mut escaped = false;
71 for ch in value_chars.chars() {
72 value_len += ch.len_utf8();
73 if escaped {
74 escaped = false;
75 } else if ch == '\\' {
76 escaped = true;
77 } else if ch == '"' {
78 break;
79 }
80 }
81 let abs_value_close = abs_value_open + 1 + value_len; let replacement = format!("\"{}\":\"[REDACTED]\"", field);
83 result.replace_range(search_from + key_pos..abs_value_close, &replacement);
84 search_from = search_from + key_pos + replacement.len();
86 }
87 }
88 result
89}
90
91#[derive(Debug, Clone, Deserialize)]
92pub struct DeviceCodeResponse {
93 pub device_code: String,
94 pub user_code: String,
95 pub verification_uri: String,
96 pub expires_in: u64,
97 pub interval: u64,
98}
99
100#[derive(Debug, Clone)]
101pub struct TokenResponse {
102 pub access_token: String,
103 pub refresh_token: String,
104 pub id_token: String,
105 pub expires_in: u64,
106 pub scope: String,
107}
108
109#[derive(Debug, Clone)]
112pub struct IdClaims {
113 pub oid: String,
114 pub tid: String,
115 pub preferred_username: String,
116 pub name: String,
117}
118
119#[derive(Deserialize)]
120struct RawTokenSuccess {
121 access_token: String,
122 refresh_token: Option<String>,
123 id_token: Option<String>,
124 expires_in: u64,
125 scope: Option<String>,
126}
127
128pub async fn request_device_code(
129 client: &reqwest::Client,
130 login_endpoint: &str,
131 tenant: &str,
132 client_id: &str,
133 scope: &str,
134) -> Result<DeviceCodeResponse> {
135 let url = format!("{login_endpoint}/{tenant}/oauth2/v2.0/devicecode");
136 let resp = client
137 .post(&url)
138 .form(&[("client_id", client_id), ("scope", scope)])
139 .send()
140 .await?;
141 if !resp.status().is_success() {
142 let status = resp.status().as_u16();
143 let body = resp.text().await.unwrap_or_default();
144 let err_msg = match serde_json::from_str::<OAuth2Error>(&body) {
145 Ok(e) => format!(
146 "device-code request failed ({status}): {}: {}",
147 e.error,
148 redact_token_fields(e.error_description.as_deref().unwrap_or_default())
149 ),
150 Err(_) => {
151 tracing::debug!("device-code request failed ({status}): {body}");
152 format!("device-code request failed with HTTP {status} and unparseable body")
153 }
154 };
155 return Err(CliError::Auth(err_msg));
156 }
157 let parsed: DeviceCodeResponse = resp.json().await?;
158 Ok(parsed)
159}
160
161async fn send_with_retry(
168 client: &reqwest::Client,
169 url: &str,
170 form: &[(&str, &str)],
171) -> Result<(reqwest::StatusCode, String)> {
172 let mut attempt: u32 = 0;
173 loop {
174 let result = client.post(url).form(form).send().await;
175
176 match result {
177 Err(e) if attempt < MAX_TOKEN_RETRIES => {
178 let backoff = 2u64.pow(attempt);
180 tracing::debug!(
181 "token endpoint connection error (attempt {attempt}): {e}; retrying in {backoff}s"
182 );
183 sleep(Duration::from_secs(backoff)).await;
184 attempt += 1;
185 continue;
186 }
187 Err(e) => return Err(CliError::Http(format!("token endpoint: {e}"))),
188 Ok(resp) => {
189 let status = resp.status();
190 let retry_after = resp
191 .headers()
192 .get("Retry-After")
193 .and_then(|v| v.to_str().ok())
194 .and_then(util::parse_retry_after);
195 let body = resp.text().await.unwrap_or_default();
196
197 let should_retry = (status == reqwest::StatusCode::TOO_MANY_REQUESTS
199 || status == reqwest::StatusCode::REQUEST_TIMEOUT
200 || status.is_server_error())
201 && attempt < MAX_TOKEN_RETRIES;
202
203 if should_retry {
204 let secs = retry_after
205 .map(|d| d.as_secs())
206 .unwrap_or_else(|| 2u64.pow(attempt));
207 tracing::debug!(
208 "token endpoint transient error {status} (attempt {attempt}); retrying in {secs}s"
209 );
210 sleep(Duration::from_secs(secs)).await;
211 attempt += 1;
212 continue;
213 }
214
215 return Ok((status, body));
216 }
217 }
218 }
219}
220
221pub async fn poll_for_token(
222 client: &reqwest::Client,
223 login_endpoint: &str,
224 tenant: &str,
225 client_id: &str,
226 device_code: &str,
227 initial_interval: u64,
228 expires_in: u64,
229) -> Result<TokenResponse> {
230 let url = format!("{login_endpoint}/{tenant}/oauth2/v2.0/token");
231 let mut interval = initial_interval.max(1);
232 let mut elapsed: u64 = 0;
233 loop {
234 if elapsed >= expires_in {
235 return Err(CliError::Auth(
236 "device code expired before sign-in completed; try again".into(),
237 ));
238 }
239 sleep(Duration::from_secs(interval)).await;
240 elapsed = elapsed.saturating_add(interval);
241
242 let (status, body) = send_with_retry(
245 client,
246 &url,
247 &[
248 ("grant_type", "urn:ietf:params:oauth:grant-type:device_code"),
249 ("client_id", client_id),
250 ("device_code", device_code),
251 ],
252 )
253 .await?;
254
255 if status.is_success() {
256 let raw: RawTokenSuccess = serde_json::from_str(&body).map_err(|e| {
257 tracing::debug!("token response body (parse error): {body}");
258 CliError::Auth(format!("token response was not valid JSON: {e}"))
259 })?;
260 return Ok(TokenResponse {
261 access_token: raw.access_token,
262 refresh_token: raw
263 .refresh_token
264 .ok_or_else(|| CliError::Auth("no refresh_token returned".into()))?,
265 id_token: raw.id_token.ok_or_else(|| {
266 CliError::Auth("no id_token returned (need 'openid' scope)".into())
267 })?,
268 expires_in: raw.expires_in,
269 scope: raw.scope.unwrap_or_default(),
270 });
271 }
272
273 let parsed: std::result::Result<OAuth2Error, _> = serde_json::from_str(&body);
275 match parsed {
276 Ok(err) => match err.error.as_str() {
277 "authorization_pending" | "bad_verification_code" => {}
278 "slow_down" => {
279 interval = interval.saturating_add(5);
280 }
281 "authorization_declined" => {
282 return Err(CliError::Auth("user declined the sign-in request".into()));
283 }
284 "expired_token" => {
285 return Err(CliError::Auth(
286 "device code expired before sign-in completed; try again".into(),
287 ));
288 }
289 "access_denied" => {
290 if err
291 .error_description
292 .as_deref()
293 .unwrap_or("")
294 .contains("AADSTS65001")
295 {
296 return Err(CliError::Auth(
297 "admin consent required for this app in your tenant; \
298 ask your IT admin to grant consent for sharepoint-cli, \
299 then try again. Details: AADSTS65001"
300 .into(),
301 ));
302 }
303 return Err(CliError::Auth(format!(
304 "access denied: {}",
305 redact_token_fields(err.error_description.as_deref().unwrap_or_default())
306 )));
307 }
308 other => {
309 return Err(CliError::Auth(format!(
310 "device-code polling failed: {other}: {}",
311 redact_token_fields(err.error_description.as_deref().unwrap_or_default())
312 )));
313 }
314 },
315 Err(_) => {
316 tracing::debug!(
317 "device-code polling failed ({status}) with unparseable body: {body}"
318 );
319 return Err(CliError::Auth(format!(
320 "token endpoint returned HTTP {status} with unparseable body"
321 )));
322 }
323 }
324 }
325}
326
327pub async fn refresh(
328 client: &reqwest::Client,
329 login_endpoint: &str,
330 tenant: &str,
331 client_id: &str,
332 refresh_token: &str,
333 scope: &str,
334) -> Result<TokenResponse> {
335 let url = format!("{login_endpoint}/{tenant}/oauth2/v2.0/token");
336 let (status, body) = send_with_retry(
337 client,
338 &url,
339 &[
340 ("grant_type", "refresh_token"),
341 ("client_id", client_id),
342 ("refresh_token", refresh_token),
343 ("scope", scope),
344 ],
345 )
346 .await?;
347 if status.is_success() {
348 let raw: RawTokenSuccess = serde_json::from_str(&body).map_err(|e| {
349 tracing::debug!("refresh response body (parse error): {body}");
350 CliError::Auth(format!("refresh response was not valid JSON: {e}"))
351 })?;
352 return Ok(TokenResponse {
353 access_token: raw.access_token,
354 refresh_token: raw
355 .refresh_token
356 .unwrap_or_else(|| refresh_token.to_string()),
357 id_token: raw.id_token.unwrap_or_default(),
358 expires_in: raw.expires_in,
359 scope: raw.scope.unwrap_or_default(),
360 });
361 }
362 let parsed: std::result::Result<OAuth2Error, _> = serde_json::from_str(&body);
363 match parsed {
364 Ok(err) if err.error == "invalid_grant" => Err(CliError::Auth(
365 "refresh token is no longer valid; run `sharepoint auth login`".into(),
366 )),
367 Ok(err) => Err(CliError::Auth(format!(
368 "refresh failed: {}: {}",
369 err.error,
370 redact_token_fields(err.error_description.as_deref().unwrap_or_default())
371 ))),
372 Err(_) => {
373 tracing::debug!("refresh failed ({status}) with unparseable body: {body}");
374 Err(CliError::Auth(format!(
375 "token endpoint returned HTTP {status} with unparseable body"
376 )))
377 }
378 }
379}
380
381pub fn decode_id_token(id_token: &str) -> Result<IdClaims> {
384 let mid = id_token
385 .split('.')
386 .nth(1)
387 .ok_or_else(|| CliError::Auth("id_token has no payload segment".into()))?;
388 let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
389 .decode(mid)
390 .map_err(|e| CliError::Auth(format!("id_token base64 decode: {e}")))?;
391 let json: serde_json::Value = serde_json::from_slice(&bytes)
392 .map_err(|e| CliError::Auth(format!("id_token JSON decode: {e}")))?;
393 let oid = json
394 .get("oid")
395 .and_then(|v| v.as_str())
396 .ok_or_else(|| CliError::Auth("id_token missing 'oid' claim".into()))?;
397 let tid = json
398 .get("tid")
399 .and_then(|v| v.as_str())
400 .ok_or_else(|| CliError::Auth("id_token missing 'tid' claim".into()))?;
401 let preferred_username = json
402 .get("preferred_username")
403 .and_then(|v| v.as_str())
404 .unwrap_or("")
405 .to_string();
406 let name = json
407 .get("name")
408 .and_then(|v| v.as_str())
409 .unwrap_or("")
410 .to_string();
411 Ok(IdClaims {
412 oid: oid.into(),
413 tid: tid.into(),
414 preferred_username,
415 name,
416 })
417}
418
419pub fn default_scope(read_only: bool) -> &'static str {
421 if read_only {
422 "openid profile offline_access User.Read Files.Read.All Sites.Read.All"
423 } else {
424 "openid profile offline_access User.Read Files.ReadWrite.All Sites.Read.All"
425 }
426}
427
428#[cfg(test)]
429mod tests {
430 use super::*;
431 use base64::Engine;
432
433 fn make_id_token(payload: &serde_json::Value) -> String {
434 let header = "{}";
435 let header_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(header);
436 let body_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD
437 .encode(serde_json::to_vec(payload).unwrap());
438 format!("{header_b64}.{body_b64}.sig")
439 }
440
441 #[test]
442 fn decode_id_token_extracts_required_claims() {
443 let token = make_id_token(&serde_json::json!({
444 "oid": "OID-123",
445 "tid": "TID-456",
446 "preferred_username": "alice@contoso.com",
447 "name": "Alice"
448 }));
449 let claims = decode_id_token(&token).unwrap();
450 assert_eq!(claims.oid, "OID-123");
451 assert_eq!(claims.tid, "TID-456");
452 assert_eq!(claims.preferred_username, "alice@contoso.com");
453 assert_eq!(claims.name, "Alice");
454 }
455
456 #[test]
457 fn decode_id_token_errors_when_oid_missing() {
458 let token = make_id_token(&serde_json::json!({"tid": "T"}));
459 assert!(decode_id_token(&token).is_err());
460 }
461
462 #[test]
463 fn default_scope_includes_files_readwrite_when_not_readonly() {
464 assert!(default_scope(false).contains("Files.ReadWrite.All"));
465 assert!(!default_scope(false).contains("Files.Read.All "));
466 }
467
468 #[test]
469 fn default_scope_uses_files_read_when_readonly() {
470 assert!(default_scope(true).contains("Files.Read.All"));
471 assert!(!default_scope(true).contains("Files.ReadWrite.All"));
472 }
473
474 #[test]
475 fn redact_token_fields_removes_access_token_value() {
476 let body = r#"{"error":"invalid_client","access_token":"SECRET-TOKEN-VALUE","foo":"bar"}"#;
477 let redacted = redact_token_fields(body);
478 assert!(
479 !redacted.contains("SECRET-TOKEN-VALUE"),
480 "access_token value must not appear in redacted output: {redacted}"
481 );
482 assert!(
483 redacted.contains("[REDACTED]"),
484 "redacted marker must appear: {redacted}"
485 );
486 assert!(
488 redacted.contains("invalid_client"),
489 "error field must survive: {redacted}"
490 );
491 }
492
493 #[test]
494 fn redact_token_fields_handles_refresh_and_id_token() {
495 let body = r#"{"refresh_token":"RT-SECRET","id_token":"IT-SECRET","ok":"keep"}"#;
496 let redacted = redact_token_fields(body);
497 assert!(!redacted.contains("RT-SECRET"));
498 assert!(!redacted.contains("IT-SECRET"));
499 assert!(redacted.contains("keep"));
500 assert_eq!(redacted.matches("[REDACTED]").count(), 2);
501 }
502
503 #[test]
504 fn redact_token_fields_is_noop_when_no_token_fields_present() {
505 let s = r#"{"error":"server_error","error_description":"oops"}"#;
506 assert_eq!(redact_token_fields(s), s);
507 }
508
509 #[tokio::test]
510 async fn poll_for_token_unparseable_5xx_does_not_leak_body() {
511 use wiremock::matchers::{method, path};
512 use wiremock::{Mock, MockServer, ResponseTemplate};
513
514 let server = MockServer::start().await;
515 Mock::given(method("POST"))
516 .and(path("/tenant/oauth2/v2.0/token"))
517 .respond_with(
518 ResponseTemplate::new(503)
519 .set_body_string(r#"{"access_token":"SHOULD_NOT_APPEAR","x":"y"}"#),
520 )
521 .mount(&server)
522 .await;
523
524 let client = reqwest::Client::new();
525 let err = poll_for_token(
526 &client,
527 &server.uri(),
528 "tenant",
529 "client-id",
530 "DEV-CODE",
531 1,
532 10,
533 )
534 .await
535 .unwrap_err();
536
537 let msg = err.to_string();
538 assert!(
539 !msg.contains("SHOULD_NOT_APPEAR"),
540 "raw body must not appear in error: {msg}"
541 );
542 assert!(
543 msg.contains("503") || msg.contains("unparseable"),
544 "error must mention status or unparseable: {msg}"
545 );
546 }
547}