1use 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
54pub struct AuthSuccess {
56 pub tokens: TokenSet,
57 pub id_token: Option<String>,
60}
61
62pub 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 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
117pub(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
167pub(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 .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
208fn 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 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 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
372fn 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("&"),
381 '<' => acc.push_str("<"),
382 '>' => acc.push_str(">"),
383 '"' => acc.push_str("""),
384 '\'' => acc.push_str("'"),
385 _ => acc.push(c),
386 }
387 acc
388 })
389}
390
391pub(crate) fn make_code_verifier() -> String {
398 let mut rng = rng();
399 (0..64).map(|_| rng.sample(Alphanumeric) as char).collect()
400}
401
402pub(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
410pub(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#[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(¶ms).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 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 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("<script>alert(1)</script>"));
544 assert!(page.contains("quote" apostrophe' amp&"));
545 }
546}