1use 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
55pub struct AuthSuccess {
57 pub tokens: TokenSet,
58 pub id_token: Option<String>,
61}
62
63pub 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 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 .append_pair("prompt", "select_account");
152 url.into()
153}
154
155struct CallbackParams {
156 code: String,
157 state: String,
158}
159
160async fn wait_for_callback(listener: TcpListener) -> Result<CallbackParams, ClientError> {
163 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 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 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
289fn make_code_verifier() -> String {
296 let mut rng = thread_rng();
297 (0..64).map(|_| rng.sample(Alphanumeric) as char).collect()
298}
299
300fn 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
308fn 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#[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(¶ms).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 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 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}