Skip to main content

pidge_client/auth/
browser_flow.rs

1//! OAuth 2.0 authorization code + PKCE flow with a one-shot local HTTP server.
2//!
3//! Why this and not device-code:
4//!
5//! Device-code flow works perfectly for work/school M365 accounts, but personal
6//! Microsoft accounts (live.com / outlook.com / hotmail.com) consistently hit
7//! `invalid_request: response_type missing` errors deep inside Microsoft's MSA
8//! pipeline. After significant debugging, and trying multiple redirect-URI
9//! shapes (nativeclient, http://localhost, urn:ietf:wg:oauth:2.0:oob), the
10//! conclusion was that personal MSA's device-code support is fragile in ways
11//! that can't be papered over via app-registration tweaks.
12//!
13//! The auth-code + PKCE + local-server flow is what modern MSAL libraries use
14//! and what Microsoft itself recommends for desktop / CLI apps. Both account
15//! types route through the same `/oauth2/v2.0/authorize` endpoint and the
16//! same redirect-URI plumbing, so M365 and MSA behave identically.
17//!
18//! Flow:
19//!
20//! 1. Bind a `TcpListener` on `127.0.0.1:0`; the OS picks a free port.
21//! 2. Generate a 64-char random `code_verifier`, derive `code_challenge =
22//!    base64url(SHA256(code_verifier))`, and a random `state` for CSRF.
23//! 3. Open the user's browser to
24//!    `https://login.microsoftonline.com/common/oauth2/v2.0/authorize?
25//!     client_id=…&response_type=code&redirect_uri=http://localhost:{port}&
26//!     scope=…&code_challenge=…&code_challenge_method=S256&state=…`
27//! 4. The user signs in; Microsoft redirects the browser back to
28//!    `http://localhost:{port}/?code=…&state=…`.
29//! 5. Our local listener accepts exactly one connection, reads the first
30//!    request line, extracts the query params, writes a friendly HTML
31//!    response, and closes.
32//! 6. We POST `/oauth2/v2.0/token` with the auth code, code_verifier, and
33//!    redirect_uri to exchange for an access + refresh token.
34//!
35//! The local server is single-shot: one connection, one response, then close.
36//! No port collisions because we let the OS pick; no listening process left
37//! behind; no firewall surprises because nothing outside the loopback
38//! interface can reach it.
39
40use std::time::Duration;
41
42use base64::Engine;
43use base64::engine::general_purpose::URL_SAFE_NO_PAD;
44use rand::distr::Alphanumeric;
45use rand::{Rng, RngExt, rng};
46use serde::Deserialize;
47use sha2::{Digest, Sha256};
48use tokio::io::{AsyncReadExt, AsyncWriteExt};
49use tokio::net::TcpListener;
50
51use crate::auth::tokens::TokenSet;
52use crate::error::ClientError;
53
54/// The result of a successful sign-in.
55pub struct AuthSuccess {
56    pub tokens: TokenSet,
57    /// The raw `id_token`, if Microsoft returned one. Caller uses this to
58    /// extract the tenant ID via the existing `jwt::extract_tenant_id`.
59    pub id_token: Option<String>,
60}
61
62/// Run the full browser flow to completion. Returns the access+refresh
63/// tokens and (optionally) the id_token for tenant extraction.
64///
65/// `on_open` is called once we know the authorize URL; the caller is
66/// expected to print it to the user and best-effort spawn their browser.
67pub async fn run<F: FnOnce(&str)>(
68    http: &reqwest::Client,
69    authority_base: &str,
70    client_id: &str,
71    scope: &str,
72    on_open: F,
73) -> Result<AuthSuccess, ClientError> {
74    // Bind first so we know which port to embed in the redirect URI.
75    let listener = TcpListener::bind("127.0.0.1:0")
76        .await
77        .map_err(ClientError::Io)?;
78    let port = listener.local_addr().map_err(ClientError::Io)?.port();
79    let redirect_uri = format!("http://localhost:{port}");
80
81    let verifier = make_code_verifier();
82    let challenge = make_code_challenge(&verifier);
83    let state = make_random(32);
84
85    let authorize_url = build_authorize_url(
86        authority_base,
87        client_id,
88        &redirect_uri,
89        scope,
90        &challenge,
91        &state,
92    );
93    on_open(&authorize_url);
94
95    let CallbackParams {
96        code,
97        state: returned_state,
98    } = wait_for_callback(listener, Duration::from_secs(300), "Microsoft sign-in").await?;
99    if returned_state != state {
100        return Err(ClientError::Graph {
101            status: 400,
102            message: "OAuth state mismatch: possible CSRF or stale request".to_string(),
103        });
104    }
105
106    exchange_code_to_success(
107        http,
108        authority_base,
109        client_id,
110        &code,
111        &verifier,
112        &redirect_uri,
113    )
114    .await
115}
116
117/// Redeem an authorization code and shape the response as an [`AuthSuccess`].
118/// Shared by the local browser flow and hosted callbacks.
119pub(crate) async fn exchange_code_to_success(
120    http: &reqwest::Client,
121    authority_base: &str,
122    client_id: &str,
123    code: &str,
124    code_verifier: &str,
125    redirect_uri: &str,
126) -> Result<AuthSuccess, ClientError> {
127    let tokens_response = exchange_code(
128        http,
129        authority_base,
130        client_id,
131        code,
132        code_verifier,
133        redirect_uri,
134    )
135    .await?;
136
137    Ok(AuthSuccess {
138        tokens: TokenSet {
139            access_token: tokens_response.access_token,
140            refresh_token: tokens_response.refresh_token.unwrap_or_default(),
141            expires_at: chrono::Utc::now()
142                + chrono::Duration::seconds(tokens_response.expires_in.unwrap_or(3600)),
143        },
144        id_token: tokens_response.id_token,
145    })
146}
147
148pub(crate) fn build_authorize_url(
149    authority_base: &str,
150    client_id: &str,
151    redirect_uri: &str,
152    scope: &str,
153    challenge: &str,
154    state: &str,
155) -> String {
156    build_authorize_url_with_hint(
157        authority_base,
158        client_id,
159        redirect_uri,
160        scope,
161        challenge,
162        state,
163        None,
164    )
165}
166
167/// [`build_authorize_url`] with an optional `login_hint`: the address
168/// Microsoft should preselect in its account picker, for flows where the
169/// caller already knows which mailbox the user is expected to sign in with.
170pub(crate) fn build_authorize_url_with_hint(
171    authority_base: &str,
172    client_id: &str,
173    redirect_uri: &str,
174    scope: &str,
175    challenge: &str,
176    state: &str,
177    login_hint: Option<&str>,
178) -> String {
179    let mut url = url::Url::parse(&format!("{authority_base}/oauth2/v2.0/authorize"))
180        .expect("authority_base is a valid URL");
181    {
182        let mut q = url.query_pairs_mut();
183        q.append_pair("client_id", client_id)
184            .append_pair("response_type", "code")
185            .append_pair("redirect_uri", redirect_uri)
186            .append_pair("response_mode", "query")
187            .append_pair("scope", scope)
188            .append_pair("state", state)
189            .append_pair("code_challenge", challenge)
190            .append_pair("code_challenge_method", "S256")
191            // `prompt=select_account` forces Microsoft to show the account picker
192            // even if the user is already signed in to *some* account. This is
193            // what stops "browser is already signed in to my M365 account so
194            // pidge auto-grabs that one when I wanted my live.com account".
195            .append_pair("prompt", "select_account");
196        if let Some(hint) = login_hint.map(str::trim).filter(|h| !h.is_empty()) {
197            q.append_pair("login_hint", hint);
198        }
199    }
200    url.into()
201}
202
203pub(crate) struct CallbackParams {
204    pub(crate) code: String,
205    pub(crate) state: String,
206}
207
208/// Accept a single connection on the listener, parse the request line for
209/// query parameters, write a success/error response, close. Single-shot.
210///
211/// `timeout` bounds how long we'll wait for the browser to hit the callback,
212/// covering both the initial connection *and* the request read that follows,
213/// so a client that opens the socket and never sends a request line still
214/// hits the deadline rather than hanging forever. The Microsoft flow uses 5
215/// minutes (matching Microsoft's own OAuth code TTL); the MCP sign-in flow
216/// (see `crate::mcp::oauth`) uses a longer window since it's a separate,
217/// often manually-triggered, sign-in step.
218///
219/// `label` identifies who this callback is for (e.g. `"Microsoft sign-in"`
220/// or `"MCP sign-in"`); it prefixes the error message when the remote party
221/// reports `error`/`error_description`, so an MCP-server error isn't
222/// misattributed to Microsoft.
223/// A timeout for the user to read: whole minutes as "5 min", else seconds.
224fn timeout_label(timeout: Duration) -> String {
225    let secs = timeout.as_secs();
226    if secs >= 60 && secs.is_multiple_of(60) {
227        format!("{} min", secs / 60)
228    } else {
229        format!("{secs}s")
230    }
231}
232
233pub(crate) async fn wait_for_callback(
234    listener: TcpListener,
235    timeout: Duration,
236    label: &str,
237) -> Result<CallbackParams, ClientError> {
238    let deadline = tokio::time::Instant::now() + timeout;
239    let timed_out = || ClientError::Graph {
240        status: 408,
241        message: format!(
242            "timed out waiting for browser sign-in ({})",
243            timeout_label(timeout)
244        ),
245    };
246
247    let (mut stream, _) = tokio::time::timeout_at(deadline, listener.accept())
248        .await
249        .map_err(|_| timed_out())?
250        .map_err(ClientError::Io)?;
251
252    // We only need the first ~1KB to parse the request line "GET /?...".
253    let mut buf = [0u8; 2048];
254    let n = tokio::time::timeout_at(deadline, stream.read(&mut buf))
255        .await
256        .map_err(|_| timed_out())?
257        .map_err(ClientError::Io)?;
258    let request = std::str::from_utf8(&buf[..n]).unwrap_or("");
259    let first_line = request.lines().next().unwrap_or("");
260    let path_and_query =
261        first_line
262            .split_whitespace()
263            .nth(1)
264            .ok_or_else(|| ClientError::Graph {
265                status: 400,
266                message: "malformed browser callback request".to_string(),
267            })?;
268
269    // Parse "/?code=...&state=..." (or "/?error=...&error_description=...").
270    let query_start = path_and_query.find('?').unwrap_or(path_and_query.len());
271    let query = &path_and_query[query_start.saturating_add(1)..];
272    let pairs: Vec<(String, String)> = url::form_urlencoded::parse(query.as_bytes())
273        .map(|(k, v)| (k.into_owned(), v.into_owned()))
274        .collect();
275
276    let mut code: Option<String> = None;
277    let mut state: Option<String> = None;
278    let mut err: Option<String> = None;
279    let mut err_description: Option<String> = None;
280    for (k, v) in pairs {
281        match k.as_str() {
282            "code" => code = Some(v),
283            "state" => state = Some(v),
284            "error" => err = Some(v),
285            "error_description" => err_description = Some(v),
286            _ => {}
287        }
288    }
289
290    if let Some(e) = err {
291        let detail = err_description.unwrap_or_default();
292        write_html(&mut stream, &error_page_html(&e, &detail), "Sign-in failed")
293            .await
294            .ok();
295        return Err(ClientError::Graph {
296            status: 400,
297            message: format!("{label}: {e} ({detail})"),
298        });
299    }
300
301    let code = code.ok_or_else(|| ClientError::Graph {
302        status: 400,
303        message: "browser callback missing `code` parameter".to_string(),
304    })?;
305    let state = state.unwrap_or_default();
306
307    write_html(&mut stream, SUCCESS_HTML, "Signed in")
308        .await
309        .ok();
310
311    Ok(CallbackParams { code, state })
312}
313
314async fn write_html(
315    stream: &mut tokio::net::TcpStream,
316    body: &str,
317    title: &str,
318) -> std::io::Result<()> {
319    let body_bytes = body.as_bytes();
320    let response = format!(
321        "HTTP/1.1 200 OK\r\n\
322         Content-Type: text/html; charset=utf-8\r\n\
323         Content-Length: {}\r\n\
324         Connection: close\r\n\
325         X-Title: {}\r\n\
326         \r\n",
327        body_bytes.len(),
328        title,
329    );
330    stream.write_all(response.as_bytes()).await?;
331    stream.write_all(body_bytes).await?;
332    stream.shutdown().await
333}
334
335const SUCCESS_HTML: &str = r#"<!doctype html><html><head><meta charset="utf-8"><title>Signed in</title>
336<style>
337body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", system-ui, sans-serif;
338       max-width: 480px; margin: 80px auto; text-align: center; color: #1d1d1f; }
339.check { font-size: 48px; color: #34c759; }
340h1 { font-size: 24px; margin: 16px 0 8px; }
341p { color: #6e6e73; }
342</style></head>
343<body>
344  <div class="check">✓</div>
345  <h1>Signed in to pidge</h1>
346  <p>You can close this window and return to the terminal.</p>
347</body></html>"#;
348
349fn error_page_html(err: &str, description: &str) -> String {
350    let err = html_escape(err);
351    let description = html_escape(description);
352    format!(
353        r#"<!doctype html><html><head><meta charset="utf-8"><title>Sign-in failed</title>
354<style>
355body {{ font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", system-ui, sans-serif;
356       max-width: 540px; margin: 80px auto; color: #1d1d1f; }}
357.x {{ font-size: 48px; color: #ff3b30; text-align: center; }}
358h1 {{ font-size: 22px; margin: 16px 0 8px; text-align: center; }}
359.detail {{ background: #f5f5f7; padding: 16px; border-radius: 8px; font-family: ui-monospace, monospace;
360         font-size: 13px; white-space: pre-wrap; word-break: break-word; }}
361</style></head>
362<body>
363  <div class="x">✕</div>
364  <h1>Sign-in failed</h1>
365  <p class="detail"><strong>{err}</strong>
366{description}</p>
367  <p>You can close this window. Return to the terminal for next steps.</p>
368</body></html>"#
369    )
370}
371
372/// Escape the five HTML-significant characters. `err` and `error_description`
373/// are fixed strings from pidge's own server, but the MCP flow (see
374/// `crate::mcp::oauth`) can point this callback at any server the user
375/// configures, so this page must not trust its input to be markup-free.
376fn html_escape(s: &str) -> String {
377    s.chars()
378        .fold(String::with_capacity(s.len()), |mut acc, c| {
379            match c {
380                '&' => acc.push_str("&amp;"),
381                '<' => acc.push_str("&lt;"),
382                '>' => acc.push_str("&gt;"),
383                '"' => acc.push_str("&quot;"),
384                '\'' => acc.push_str("&#39;"),
385                _ => acc.push(c),
386            }
387            acc
388        })
389}
390
391// --- PKCE helpers ----------------------------------------------------------
392
393/// RFC 7636 says the code_verifier is "a high-entropy cryptographic random
394/// STRING, using the unreserved characters … with a minimum length of 43
395/// characters and a maximum length of 128 characters." 64 alphanumerics is
396/// comfortably inside the spec and gives ~380 bits of entropy.
397pub(crate) fn make_code_verifier() -> String {
398    let mut rng = rng();
399    (0..64).map(|_| rng.sample(Alphanumeric) as char).collect()
400}
401
402/// `base64url(SHA256(code_verifier))` per RFC 7636 §4.2. URL_SAFE_NO_PAD is
403/// the exact encoding the OAuth spec requires.
404pub(crate) fn make_code_challenge(verifier: &str) -> String {
405    let mut hasher = Sha256::new();
406    hasher.update(verifier.as_bytes());
407    URL_SAFE_NO_PAD.encode(hasher.finalize())
408}
409
410/// 32 bytes of OS random → URL-safe base64. Used for the `state` CSRF nonce.
411pub(crate) fn make_random(byte_len: usize) -> String {
412    let mut buf = vec![0u8; byte_len];
413    rng().fill_bytes(&mut buf);
414    URL_SAFE_NO_PAD.encode(&buf)
415}
416
417// --- token exchange --------------------------------------------------------
418
419#[derive(Debug, Deserialize)]
420struct TokenResponse {
421    access_token: String,
422    refresh_token: Option<String>,
423    expires_in: Option<i64>,
424    id_token: Option<String>,
425}
426
427async fn exchange_code(
428    http: &reqwest::Client,
429    authority_base: &str,
430    client_id: &str,
431    code: &str,
432    code_verifier: &str,
433    redirect_uri: &str,
434) -> Result<TokenResponse, ClientError> {
435    let url = format!("{authority_base}/oauth2/v2.0/token");
436    let params = [
437        ("client_id", client_id),
438        ("grant_type", "authorization_code"),
439        ("code", code),
440        ("code_verifier", code_verifier),
441        ("redirect_uri", redirect_uri),
442    ];
443    let resp = http.post(&url).form(&params).send().await?;
444    let status = resp.status();
445    if !status.is_success() {
446        let text = resp.text().await.unwrap_or_default();
447        return Err(ClientError::Graph {
448            status: status.as_u16(),
449            message: text,
450        });
451    }
452    Ok(resp.json().await?)
453}
454
455#[cfg(test)]
456mod tests {
457    use super::*;
458
459    #[test]
460    fn timeout_label_uses_minutes_when_whole() {
461        assert_eq!(timeout_label(Duration::from_secs(300)), "5 min");
462        assert_eq!(timeout_label(Duration::from_secs(600)), "10 min");
463        assert_eq!(timeout_label(Duration::from_secs(45)), "45s");
464    }
465
466    #[test]
467    fn code_verifier_is_64_alphanumerics() {
468        let v = make_code_verifier();
469        assert_eq!(v.len(), 64);
470        assert!(v.chars().all(|c| c.is_ascii_alphanumeric()));
471    }
472
473    #[test]
474    fn challenge_matches_rfc_7636_example() {
475        // RFC 7636 §4 example:
476        //   verifier  = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"
477        //   challenge = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
478        let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
479        assert_eq!(
480            make_code_challenge(verifier),
481            "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
482        );
483    }
484
485    #[test]
486    fn authorize_url_contains_required_params() {
487        let url = build_authorize_url(
488            "https://login.microsoftonline.com/common",
489            "client-id-here",
490            "http://localhost:47821",
491            "User.Read offline_access",
492            "challenge-here",
493            "state-here",
494        );
495        assert!(url.contains("client_id=client-id-here"));
496        assert!(url.contains("response_type=code"));
497        assert!(url.contains("code_challenge=challenge-here"));
498        assert!(url.contains("code_challenge_method=S256"));
499        assert!(url.contains("state=state-here"));
500        // redirect_uri is percent-encoded inside the query string.
501        assert!(url.contains("redirect_uri=http%3A%2F%2Flocalhost%3A47821"));
502        assert!(url.contains("prompt=select_account"));
503        assert!(!url.contains("login_hint"));
504    }
505
506    #[test]
507    fn authorize_url_carries_the_login_hint_when_given() {
508        let url = build_authorize_url_with_hint(
509            "https://login.microsoftonline.com/common",
510            "client-id-here",
511            "http://localhost:47821",
512            "scope-here",
513            "challenge-here",
514            "state-here",
515            Some("jane.doe@example.com"),
516        );
517        assert!(url.contains("login_hint=jane.doe%40example.com"), "{url}");
518        assert!(url.contains("prompt=select_account"));
519        let blank = build_authorize_url_with_hint(
520            "https://login.microsoftonline.com/common",
521            "c",
522            "http://localhost:1",
523            "s",
524            "ch",
525            "st",
526            Some("  "),
527        );
528        assert!(!blank.contains("login_hint"));
529    }
530
531    #[test]
532    fn random_state_is_unique_per_call() {
533        let a = make_random(32);
534        let b = make_random(32);
535        assert_ne!(a, b);
536        assert_eq!(URL_SAFE_NO_PAD.decode(&a).unwrap().len(), 32);
537    }
538
539    #[test]
540    fn error_page_html_escapes_html_special_characters() {
541        let page = error_page_html("<script>alert(1)</script>", "quote\" apostrophe' amp&");
542        assert!(!page.contains("<script>"));
543        assert!(page.contains("&lt;script&gt;alert(1)&lt;/script&gt;"));
544        assert!(page.contains("quote&quot; apostrophe&#39; amp&amp;"));
545    }
546}