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