Skip to main content

sharepoint_cli/auth/
device_code.rs

1//! Device-code flow against `login.microsoftonline.com/<tenant>/oauth2/v2.0/`.
2//!
3//! Polling state machine handles all the cases the spec calls out:
4//! - 200 OK → success
5//! - 400 authorization_pending → keep polling at the same interval
6//! - 400 slow_down → bump interval by +5s
7//! - 400 bad_verification_code → keep polling (transient)
8//! - 400 authorization_declined / expired_token / access_denied → terminal failure
9//! - any other 4xx/5xx → terminal failure with structured error (body never leaked)
10//!
11//! Polling budget tracks scheduled sleep time only; real wall clock can exceed
12//! `expires_in` if requests are slow. The server's `expired_token` response is
13//! the authoritative cap.
14
15use std::time::Duration;
16
17use base64::Engine;
18use serde::Deserialize;
19use tokio::time::sleep;
20
21use crate::error::{CliError, Result};
22use crate::util;
23
24/// Maximum number of automatic retries on transient token-endpoint errors
25/// (5xx, 429, connection failures). Does not apply to the normal
26/// authorization_pending / slow_down state-machine cycles.
27const MAX_TOKEN_RETRIES: u32 = 3;
28
29/// Structured OAuth2 error response from the token endpoint.
30#[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
37/// Strip token-shaped substrings from any string destined for an error message.
38///
39/// Prevents OAuth token values from leaking into `CliError` messages when
40/// unexpected fields appear in malformed responses. Scans for JSON key-value
41/// pairs where the key is a known token field name and replaces the value with
42/// `[REDACTED]`.
43pub 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        // Replace `"<field>":"<anything>"` patterns (with optional whitespace).
48        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            // Skip whitespace after the key.
53            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            // Find where the value string ends (first unescaped closing quote).
66            let value_start_offset = result[abs_colon_end..].len() - after_colon.len();
67            let abs_value_open = abs_colon_end + value_start_offset; // points at opening "
68            let value_chars = &result[abs_value_open + 1..]; // skip opening "
69            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; // points just past closing "
82            let replacement = format!("\"{}\":\"[REDACTED]\"", field);
83            result.replace_range(search_from + key_pos..abs_value_close, &replacement);
84            // Advance past the replacement to avoid re-processing.
85            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/// Identity claims we extract from the id_token (`oid`, `tid`,
110/// `preferred_username`, `name`).
111#[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
161/// POST `url` with the given form fields, retrying on 5xx / 429 / connection
162/// failures up to [`MAX_TOKEN_RETRIES`] times with exponential backoff.
163/// Honors `Retry-After` response headers. Returns the final `(status, body)`.
164///
165/// Terminal 4xx responses (except 408 Request Timeout and 429) are returned
166/// immediately without retry — they indicate a hard auth failure.
167async 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                // Connection-level failure (timeout, reset, DNS): retry.
179                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                // Retry on 429 and 5xx; return everything else immediately.
198                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        // Inner retry loop handles transient errors (5xx, 429, connection
243        // failures) without consuming device-code expiry budget.
244        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        // Non-200: classify the OAuth error code.
274        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
381/// Decode the middle segment of a JWT (no signature verification — we trust
382/// the channel the token came over, like every other MSAL-style client).
383pub 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
419/// Build the full scope string we request in v0.1.
420pub 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        // Non-token fields must be preserved.
487        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}