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` — 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::RngCore;
45use rand::distributions::Alphanumeric;
46use rand::{Rng, thread_rng};
47use serde::Deserialize;
48use sha2::{Digest, Sha256};
49use tokio::io::{AsyncReadExt, AsyncWriteExt};
50use tokio::net::TcpListener;
51
52use crate::auth::tokens::TokenSet;
53use crate::error::ClientError;
54
55/// The result of a successful sign-in.
56pub struct AuthSuccess {
57    pub tokens: TokenSet,
58    /// The raw `id_token`, if Microsoft returned one. Caller uses this to
59    /// extract the tenant ID via the existing `jwt::extract_tenant_id`.
60    pub id_token: Option<String>,
61}
62
63/// Run the full browser flow to completion. Returns the access+refresh
64/// tokens and (optionally) the id_token for tenant extraction.
65///
66/// `on_open` is called once we know the authorize URL — the caller is
67/// expected to print it to the user and best-effort spawn their browser.
68pub async fn run<F: FnOnce(&str)>(
69    http: &reqwest::Client,
70    authority_base: &str,
71    client_id: &str,
72    scope: &str,
73    on_open: F,
74) -> Result<AuthSuccess, ClientError> {
75    // Bind first so we know which port to embed in the redirect URI.
76    let listener = TcpListener::bind("127.0.0.1:0")
77        .await
78        .map_err(ClientError::Io)?;
79    let port = listener.local_addr().map_err(ClientError::Io)?.port();
80    let redirect_uri = format!("http://localhost:{port}");
81
82    let verifier = make_code_verifier();
83    let challenge = make_code_challenge(&verifier);
84    let state = make_random(32);
85
86    let authorize_url = build_authorize_url(
87        authority_base,
88        client_id,
89        &redirect_uri,
90        scope,
91        &challenge,
92        &state,
93    );
94    on_open(&authorize_url);
95
96    let CallbackParams {
97        code,
98        state: returned_state,
99    } = wait_for_callback(listener).await?;
100    if returned_state != state {
101        return Err(ClientError::Graph {
102            status: 400,
103            message: "OAuth state mismatch — possible CSRF or stale request".to_string(),
104        });
105    }
106
107    let tokens_response = exchange_code(
108        http,
109        authority_base,
110        client_id,
111        &code,
112        &verifier,
113        &redirect_uri,
114    )
115    .await?;
116
117    Ok(AuthSuccess {
118        tokens: TokenSet {
119            access_token: tokens_response.access_token,
120            refresh_token: tokens_response.refresh_token.unwrap_or_default(),
121            expires_at: chrono::Utc::now()
122                + chrono::Duration::seconds(tokens_response.expires_in.unwrap_or(3600)),
123        },
124        id_token: tokens_response.id_token,
125    })
126}
127
128fn build_authorize_url(
129    authority_base: &str,
130    client_id: &str,
131    redirect_uri: &str,
132    scope: &str,
133    challenge: &str,
134    state: &str,
135) -> String {
136    let mut url = url::Url::parse(&format!("{authority_base}/oauth2/v2.0/authorize"))
137        .expect("authority_base is a valid URL");
138    url.query_pairs_mut()
139        .append_pair("client_id", client_id)
140        .append_pair("response_type", "code")
141        .append_pair("redirect_uri", redirect_uri)
142        .append_pair("response_mode", "query")
143        .append_pair("scope", scope)
144        .append_pair("state", state)
145        .append_pair("code_challenge", challenge)
146        .append_pair("code_challenge_method", "S256")
147        // `prompt=select_account` forces Microsoft to show the account picker
148        // even if the user is already signed in to *some* account — this is
149        // what stops "browser is already signed in to my M365 account so
150        // pidge auto-grabs that one when I wanted my live.com account".
151        .append_pair("prompt", "select_account");
152    url.into()
153}
154
155struct CallbackParams {
156    code: String,
157    state: String,
158}
159
160/// Accept a single connection on the listener, parse the request line for
161/// query parameters, write a success/error response, close. Single-shot.
162async fn wait_for_callback(listener: TcpListener) -> Result<CallbackParams, ClientError> {
163    // Generous timeout: users might take a minute or two to authenticate,
164    // especially on MFA. 5 minutes matches Microsoft's own OAuth code TTL.
165    let accept = listener.accept();
166    let (mut stream, _) = tokio::time::timeout(Duration::from_secs(300), accept)
167        .await
168        .map_err(|_| ClientError::Graph {
169            status: 408,
170            message: "timed out waiting for browser sign-in (5 min)".to_string(),
171        })?
172        .map_err(ClientError::Io)?;
173
174    // We only need the first ~1KB to parse the request line "GET /?...".
175    let mut buf = [0u8; 2048];
176    let n = stream.read(&mut buf).await.map_err(ClientError::Io)?;
177    let request = std::str::from_utf8(&buf[..n]).unwrap_or("");
178    let first_line = request.lines().next().unwrap_or("");
179    let path_and_query =
180        first_line
181            .split_whitespace()
182            .nth(1)
183            .ok_or_else(|| ClientError::Graph {
184                status: 400,
185                message: "malformed browser callback request".to_string(),
186            })?;
187
188    // Parse "/?code=...&state=..." (or "/?error=...&error_description=...").
189    let query_start = path_and_query.find('?').unwrap_or(path_and_query.len());
190    let query = &path_and_query[query_start.saturating_add(1)..];
191    let pairs: Vec<(String, String)> = url::form_urlencoded::parse(query.as_bytes())
192        .map(|(k, v)| (k.into_owned(), v.into_owned()))
193        .collect();
194
195    let mut code: Option<String> = None;
196    let mut state: Option<String> = None;
197    let mut err: Option<String> = None;
198    let mut err_description: Option<String> = None;
199    for (k, v) in pairs {
200        match k.as_str() {
201            "code" => code = Some(v),
202            "state" => state = Some(v),
203            "error" => err = Some(v),
204            "error_description" => err_description = Some(v),
205            _ => {}
206        }
207    }
208
209    if let Some(e) = err {
210        let detail = err_description.unwrap_or_default();
211        write_html(&mut stream, &error_page_html(&e, &detail), "Sign-in failed")
212            .await
213            .ok();
214        return Err(ClientError::Graph {
215            status: 400,
216            message: format!("Microsoft sign-in: {e} — {detail}"),
217        });
218    }
219
220    let code = code.ok_or_else(|| ClientError::Graph {
221        status: 400,
222        message: "browser callback missing `code` parameter".to_string(),
223    })?;
224    let state = state.unwrap_or_default();
225
226    write_html(&mut stream, SUCCESS_HTML, "Signed in")
227        .await
228        .ok();
229
230    Ok(CallbackParams { code, state })
231}
232
233async fn write_html(
234    stream: &mut tokio::net::TcpStream,
235    body: &str,
236    title: &str,
237) -> std::io::Result<()> {
238    let body_bytes = body.as_bytes();
239    let response = format!(
240        "HTTP/1.1 200 OK\r\n\
241         Content-Type: text/html; charset=utf-8\r\n\
242         Content-Length: {}\r\n\
243         Connection: close\r\n\
244         X-Title: {}\r\n\
245         \r\n",
246        body_bytes.len(),
247        title,
248    );
249    stream.write_all(response.as_bytes()).await?;
250    stream.write_all(body_bytes).await?;
251    stream.shutdown().await
252}
253
254const SUCCESS_HTML: &str = r#"<!doctype html><html><head><meta charset="utf-8"><title>Signed in</title>
255<style>
256body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", system-ui, sans-serif;
257       max-width: 480px; margin: 80px auto; text-align: center; color: #1d1d1f; }
258.check { font-size: 48px; color: #34c759; }
259h1 { font-size: 24px; margin: 16px 0 8px; }
260p { color: #6e6e73; }
261</style></head>
262<body>
263  <div class="check">✓</div>
264  <h1>Signed in to pidge</h1>
265  <p>You can close this window and return to the terminal.</p>
266</body></html>"#;
267
268fn error_page_html(err: &str, description: &str) -> String {
269    format!(
270        r#"<!doctype html><html><head><meta charset="utf-8"><title>Sign-in failed</title>
271<style>
272body {{ font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", system-ui, sans-serif;
273       max-width: 540px; margin: 80px auto; color: #1d1d1f; }}
274.x {{ font-size: 48px; color: #ff3b30; text-align: center; }}
275h1 {{ font-size: 22px; margin: 16px 0 8px; text-align: center; }}
276.detail {{ background: #f5f5f7; padding: 16px; border-radius: 8px; font-family: ui-monospace, monospace;
277         font-size: 13px; white-space: pre-wrap; word-break: break-word; }}
278</style></head>
279<body>
280  <div class="x">✕</div>
281  <h1>Sign-in failed</h1>
282  <p class="detail"><strong>{err}</strong>
283{description}</p>
284  <p>You can close this window. Return to the terminal for next steps.</p>
285</body></html>"#
286    )
287}
288
289// --- PKCE helpers ----------------------------------------------------------
290
291/// RFC 7636 says the code_verifier is "a high-entropy cryptographic random
292/// STRING, using the unreserved characters … with a minimum length of 43
293/// characters and a maximum length of 128 characters." 64 alphanumerics is
294/// comfortably inside the spec and gives ~380 bits of entropy.
295fn make_code_verifier() -> String {
296    let mut rng = thread_rng();
297    (0..64).map(|_| rng.sample(Alphanumeric) as char).collect()
298}
299
300/// `base64url(SHA256(code_verifier))` per RFC 7636 §4.2. URL_SAFE_NO_PAD is
301/// the exact encoding the OAuth spec requires.
302fn make_code_challenge(verifier: &str) -> String {
303    let mut hasher = Sha256::new();
304    hasher.update(verifier.as_bytes());
305    URL_SAFE_NO_PAD.encode(hasher.finalize())
306}
307
308/// 32 bytes of OS random → URL-safe base64. Used for the `state` CSRF nonce.
309fn make_random(byte_len: usize) -> String {
310    let mut buf = vec![0u8; byte_len];
311    thread_rng().fill_bytes(&mut buf);
312    URL_SAFE_NO_PAD.encode(&buf)
313}
314
315// --- token exchange --------------------------------------------------------
316
317#[derive(Debug, Deserialize)]
318struct TokenResponse {
319    access_token: String,
320    refresh_token: Option<String>,
321    expires_in: Option<i64>,
322    id_token: Option<String>,
323}
324
325async fn exchange_code(
326    http: &reqwest::Client,
327    authority_base: &str,
328    client_id: &str,
329    code: &str,
330    code_verifier: &str,
331    redirect_uri: &str,
332) -> Result<TokenResponse, ClientError> {
333    let url = format!("{authority_base}/oauth2/v2.0/token");
334    let params = [
335        ("client_id", client_id),
336        ("grant_type", "authorization_code"),
337        ("code", code),
338        ("code_verifier", code_verifier),
339        ("redirect_uri", redirect_uri),
340    ];
341    let resp = http.post(&url).form(&params).send().await?;
342    let status = resp.status();
343    if !status.is_success() {
344        let text = resp.text().await.unwrap_or_default();
345        return Err(ClientError::Graph {
346            status: status.as_u16(),
347            message: text,
348        });
349    }
350    Ok(resp.json().await?)
351}
352
353#[cfg(test)]
354mod tests {
355    use super::*;
356
357    #[test]
358    fn code_verifier_is_64_alphanumerics() {
359        let v = make_code_verifier();
360        assert_eq!(v.len(), 64);
361        assert!(v.chars().all(|c| c.is_ascii_alphanumeric()));
362    }
363
364    #[test]
365    fn challenge_matches_rfc_7636_example() {
366        // RFC 7636 §4 example:
367        //   verifier  = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"
368        //   challenge = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
369        let verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk";
370        assert_eq!(
371            make_code_challenge(verifier),
372            "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
373        );
374    }
375
376    #[test]
377    fn authorize_url_contains_required_params() {
378        let url = build_authorize_url(
379            "https://login.microsoftonline.com/common",
380            "client-id-here",
381            "http://localhost:47821",
382            "User.Read offline_access",
383            "challenge-here",
384            "state-here",
385        );
386        assert!(url.contains("client_id=client-id-here"));
387        assert!(url.contains("response_type=code"));
388        assert!(url.contains("code_challenge=challenge-here"));
389        assert!(url.contains("code_challenge_method=S256"));
390        assert!(url.contains("state=state-here"));
391        // redirect_uri is percent-encoded inside the query string.
392        assert!(url.contains("redirect_uri=http%3A%2F%2Flocalhost%3A47821"));
393        assert!(url.contains("prompt=select_account"));
394    }
395
396    #[test]
397    fn random_state_is_unique_per_call() {
398        let a = make_random(32);
399        let b = make_random(32);
400        assert_ne!(a, b);
401        assert_eq!(URL_SAFE_NO_PAD.decode(&a).unwrap().len(), 32);
402    }
403}