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).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 .append_pair("prompt", "select_account");
151 url.into()
152}
153
154struct CallbackParams {
155 code: String,
156 state: String,
157}
158
159async fn wait_for_callback(listener: TcpListener) -> Result<CallbackParams, ClientError> {
162 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 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 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
288fn make_code_verifier() -> String {
295 let mut rng = rng();
296 (0..64).map(|_| rng.sample(Alphanumeric) as char).collect()
297}
298
299fn 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
307fn 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#[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(¶ms).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 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 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}