Skip to main content

max / goingson

24.5 KB · 724 lines History Blame Raw
1 //! Local HTTP server for OAuth2 redirect callbacks.
2 //!
3 //! Runs a minimal HTTP server on localhost to receive the OAuth callback
4 //! after the user authorizes in their browser.
5
6 use std::io::{Read, Write};
7 use std::net::{TcpListener, TcpStream};
8 use std::sync::mpsc::{self, Receiver, Sender};
9 use std::sync::{Arc, Mutex};
10 use std::thread;
11 use std::time::Duration;
12
13 /// Escape a string for safe inclusion in HTML content.
14 fn html_escape(s: &str) -> String {
15 s.replace('&', "&")
16 .replace('<', "&lt;")
17 .replace('>', "&gt;")
18 .replace('"', "&quot;")
19 .replace('\'', "&#x27;")
20 }
21
22 /// Result of an OAuth callback.
23 #[derive(Debug, Clone)]
24 pub struct CallbackResult {
25 /// Authorization code from the OAuth provider.
26 pub code: String,
27 /// State parameter to verify CSRF.
28 pub state: String,
29 }
30
31 /// Error that occurred during OAuth callback.
32 #[derive(Debug, Clone)]
33 pub struct CallbackError {
34 /// Error code from OAuth provider.
35 pub error: String,
36 /// Human-readable error description.
37 pub error_description: Option<String>,
38 }
39
40 /// Stored callback data for polling.
41 #[derive(Debug, Clone)]
42 pub enum StoredCallback {
43 Pending,
44 Success { code: String, state: String },
45 Error { error: String, description: Option<String> },
46 }
47
48 /// A local HTTP server that handles OAuth2 redirects.
49 pub struct OAuthCallbackServer {
50 port: u16,
51 receiver: Receiver<Result<CallbackResult, CallbackError>>,
52 /// Shared callback state, surfaced to the frontend through the trusted IPC
53 /// layer rather than an unauthenticated local HTTP endpoint.
54 stored: Arc<Mutex<StoredCallback>>,
55 }
56
57 impl OAuthCallbackServer {
58 /// Starts a new callback server on a random available port.
59 ///
60 /// Returns the server and the port it's listening on.
61 pub fn start() -> Result<Self, String> {
62 // Bind to port 0 to get a random available port
63 let listener = TcpListener::bind("127.0.0.1:0")
64 .map_err(|e| format!("Failed to bind callback server: {}", e))?;
65
66 let port = listener
67 .local_addr()
68 .map_err(|e| format!("Failed to get server port: {}", e))?
69 .port();
70
71 // Set non-blocking so we can timeout
72 listener
73 .set_nonblocking(true)
74 .map_err(|e| format!("Failed to set non-blocking: {}", e))?;
75
76 let (sender, receiver) = mpsc::channel();
77
78 // Shared storage for callback result (for polling)
79 let stored = Arc::new(Mutex::new(StoredCallback::Pending));
80
81 // Spawn a thread to handle the callback
82 let stored_clone = stored.clone();
83 thread::spawn(move || {
84 Self::run_server(listener, sender, stored_clone);
85 });
86
87 Ok(Self { port, receiver, stored })
88 }
89
90 /// Returns the port this server is listening on.
91 pub fn port(&self) -> u16 {
92 self.port
93 }
94
95 /// Returns the current callback state for delivery over the trusted IPC layer.
96 ///
97 /// Used instead of an HTTP `/result` endpoint so the authorization code and
98 /// state are never readable by other local processes on the loopback port.
99 pub fn poll(&self) -> StoredCallback {
100 self.stored.lock().unwrap_or_else(|e| e.into_inner()).clone()
101 }
102
103 /// Waits for the OAuth callback with a timeout.
104 ///
105 /// # Arguments
106 /// * `timeout` - Maximum time to wait for the callback.
107 pub fn wait_for_callback(
108 &self,
109 timeout: Duration,
110 ) -> Result<Result<CallbackResult, CallbackError>, String> {
111 self.receiver
112 .recv_timeout(timeout)
113 .map_err(|e| format!("Timeout waiting for OAuth callback: {}", e))
114 }
115
116 fn run_server(
117 listener: TcpListener,
118 sender: Sender<Result<CallbackResult, CallbackError>>,
119 stored: Arc<Mutex<StoredCallback>>,
120 ) {
121 // Accept connections for up to 5 minutes
122 let deadline = std::time::Instant::now() + Duration::from_secs(300);
123 let mut callback_received = false;
124
125 while std::time::Instant::now() < deadline {
126 match listener.accept() {
127 Ok((stream, _)) => {
128 let result = Self::handle_request(stream, &stored, callback_received);
129 if let Some(callback_result) = result {
130 let _ = sender.send(callback_result);
131 callback_received = true;
132 // Continue running for a bit to serve /result requests
133 }
134 }
135 Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
136 thread::sleep(Duration::from_millis(100));
137 }
138 Err(_) => {
139 thread::sleep(Duration::from_millis(100));
140 }
141 }
142
143 // If we received a callback, keep serving for 30 more seconds for polling
144 if callback_received {
145 let poll_deadline = std::time::Instant::now() + Duration::from_secs(30);
146 while std::time::Instant::now() < poll_deadline {
147 match listener.accept() {
148 Ok((stream, _)) => {
149 Self::handle_request(stream, &stored, true);
150 }
151 Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
152 thread::sleep(Duration::from_millis(100));
153 }
154 Err(_) => {
155 thread::sleep(Duration::from_millis(100));
156 }
157 }
158 }
159 return;
160 }
161 }
162 }
163
164 fn handle_request(
165 mut stream: TcpStream,
166 stored: &Arc<Mutex<StoredCallback>>,
167 _callback_received: bool,
168 ) -> Option<Result<CallbackResult, CallbackError>> {
169 let mut buffer = [0; 16384];
170 let n = stream.read(&mut buffer).ok()?;
171 let request = String::from_utf8_lossy(&buffer[..n]);
172
173 // Parse the GET request line
174 let first_line = request.lines().next()?;
175 if !first_line.starts_with("GET ") {
176 Self::send_response(&mut stream, "405 Method Not Allowed", "text/plain", "Only GET is supported");
177 return None;
178 }
179
180 // Extract path and query string
181 let path = first_line
182 .strip_prefix("GET ")?
183 .split_whitespace()
184 .next()?;
185
186 // Parse query parameters for the callback
187 let query = path.split('?').nth(1).unwrap_or("");
188 let params: std::collections::HashMap<&str, &str> = query
189 .split('&')
190 .filter_map(|pair| {
191 let mut parts = pair.splitn(2, '=');
192 Some((parts.next()?, parts.next().unwrap_or("")))
193 })
194 .collect();
195
196 // Check for error
197 if let Some(error) = params.get("error") {
198 let error_description = params.get("error_description").map(|s| {
199 urlencoding::decode(s).unwrap_or_else(|_| s.to_string())
200 });
201
202 // Store the error
203 if let Ok(mut stored_guard) = stored.lock() {
204 *stored_guard = StoredCallback::Error {
205 error: error.to_string(),
206 description: error_description.clone(),
207 };
208 }
209
210 let safe_msg = html_escape(error_description.as_deref().unwrap_or(error));
211 Self::send_response(
212 &mut stream,
213 "200 OK",
214 "text/html; charset=utf-8",
215 &format!(
216 r#"<!DOCTYPE html>
217 <html>
218 <head><title>Authorization Failed</title></head>
219 <body style="font-family: system-ui; padding: 2rem; text-align: center;">
220 <h1 style="color: #d33;">Authorization Failed</h1>
221 <p>{}</p>
222 <p style="color: #666;">You can close this window.</p>
223 </body>
224 </html>"#,
225 safe_msg
226 ),
227 );
228
229 return Some(Err(CallbackError {
230 error: error.to_string(),
231 error_description,
232 }));
233 }
234
235 // Extract code and state, URL-decoding to handle encoded chars
236 let code = params.get("code")?;
237 let state = params.get("state")?;
238 let code = urlencoding::decode(code).unwrap_or_else(|_| code.to_string());
239 let state = urlencoding::decode(state).unwrap_or_else(|_| state.to_string());
240
241 // Store the success result
242 if let Ok(mut stored_guard) = stored.lock() {
243 *stored_guard = StoredCallback::Success {
244 code: code.clone(),
245 state: state.clone(),
246 };
247 }
248
249 Self::send_response(
250 &mut stream,
251 "200 OK",
252 "text/html; charset=utf-8",
253 r#"<!DOCTYPE html>
254 <html>
255 <head><title>Authorization Successful</title></head>
256 <body style="font-family: system-ui; padding: 2rem; text-align: center;">
257 <h1 style="color: #090;">Authorization Successful</h1>
258 <p>Your email account has been connected.</p>
259 <p style="color: #666;">You can close this window and return to GoingsOn.</p>
260 <script>setTimeout(function() { window.close(); }, 2000);</script>
261 </body>
262 </html>"#,
263 );
264
265 Some(Ok(CallbackResult {
266 code,
267 state,
268 }))
269 }
270
271 fn send_response(stream: &mut TcpStream, status: &str, content_type: &str, body: &str) {
272 // No CORS header: these pages are rendered by the user's external
273 // browser after the redirect, and the callback result is delivered to
274 // the app over IPC, so no cross-origin fetch needs to read this.
275 let response = format!(
276 "HTTP/1.1 {}\r\nContent-Type: {}\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
277 status,
278 content_type,
279 body.len(),
280 body
281 );
282 let _ = stream.write_all(response.as_bytes());
283 let _ = stream.flush();
284 }
285 }
286
287 /// Simple URL decoding with proper multi-byte UTF-8 support.
288 mod urlencoding {
289 pub fn decode(s: &str) -> Result<String, ()> {
290 let mut bytes = Vec::with_capacity(s.len());
291 let mut chars = s.chars();
292
293 while let Some(c) = chars.next() {
294 if c == '%' {
295 let hex: String = chars.by_ref().take(2).collect();
296 if hex.len() == 2
297 && let Ok(byte) = u8::from_str_radix(&hex, 16) {
298 bytes.push(byte);
299 continue;
300 }
301 return Err(());
302 } else if c == '+' {
303 bytes.push(b' ');
304 } else {
305 let mut buf = [0u8; 4];
306 let encoded = c.encode_utf8(&mut buf);
307 bytes.extend_from_slice(encoded.as_bytes());
308 }
309 }
310
311 String::from_utf8(bytes).map_err(|_| ())
312 }
313 }
314
315 #[cfg(test)]
316 mod tests {
317 use super::*;
318 use std::net::TcpStream;
319 use std::sync::Mutex;
320 use std::time::Duration;
321
322 // Serialize server integration tests to avoid port/timing interference
323 // when multiple tests spawn callback servers simultaneously.
324 // Uses unwrap_or_else to recover from a poisoned mutex (if a prior test panicked).
325 static SERVER_TEST_LOCK: Mutex<()> = Mutex::new(());
326
327 fn lock_server_tests() -> std::sync::MutexGuard<'static, ()> {
328 SERVER_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner())
329 }
330
331 // ============ URL Decoding Tests ============
332
333 #[test]
334 fn url_decode_plain_string() {
335 assert_eq!(urlencoding::decode("hello").unwrap(), "hello");
336 }
337
338 #[test]
339 fn url_decode_percent_encoded_spaces() {
340 assert_eq!(urlencoding::decode("hello%20world").unwrap(), "hello world");
341 }
342
343 #[test]
344 fn url_decode_plus_as_space() {
345 assert_eq!(urlencoding::decode("hello+world").unwrap(), "hello world");
346 }
347
348 #[test]
349 fn url_decode_special_characters() {
350 assert_eq!(urlencoding::decode("a%3Db%26c%3Dd").unwrap(), "a=b&c=d");
351 }
352
353 #[test]
354 fn url_decode_slash() {
355 assert_eq!(urlencoding::decode("%2F").unwrap(), "/");
356 }
357
358 #[test]
359 fn url_decode_mixed_encoded_and_plain() {
360 assert_eq!(
361 urlencoding::decode("access%20denied%3A+invalid+scope").unwrap(),
362 "access denied: invalid scope"
363 );
364 }
365
366 #[test]
367 fn url_decode_empty_string() {
368 assert_eq!(urlencoding::decode("").unwrap(), "");
369 }
370
371 #[test]
372 fn url_decode_truncated_percent_sequence() {
373 // "%A" has only one hex digit instead of two
374 assert!(urlencoding::decode("%A").is_err());
375 }
376
377 #[test]
378 fn url_decode_invalid_hex_after_percent() {
379 assert!(urlencoding::decode("%ZZ").is_err());
380 }
381
382 #[test]
383 fn url_decode_percent_at_end() {
384 // "%" with nothing after it
385 assert!(urlencoding::decode("hello%").is_err());
386 }
387
388 // ============ CallbackResult / CallbackError struct tests ============
389
390 #[test]
391 fn callback_result_stores_code_and_state() {
392 let result = CallbackResult {
393 code: "auth_code_123".to_string(),
394 state: "csrf_state_abc".to_string(),
395 };
396 assert_eq!(result.code, "auth_code_123");
397 assert_eq!(result.state, "csrf_state_abc");
398 }
399
400 #[test]
401 fn callback_error_with_description() {
402 let err = CallbackError {
403 error: "access_denied".to_string(),
404 error_description: Some("User denied access".to_string()),
405 };
406 assert_eq!(err.error, "access_denied");
407 assert_eq!(err.error_description.as_deref(), Some("User denied access"));
408 }
409
410 #[test]
411 fn callback_error_without_description() {
412 let err = CallbackError {
413 error: "server_error".to_string(),
414 error_description: None,
415 };
416 assert_eq!(err.error, "server_error");
417 assert!(err.error_description.is_none());
418 }
419
420 // ============ Server Integration Tests ============
421
422 /// Helper: send a raw HTTP GET request and read the full response.
423 ///
424 /// Reads headers first, extracts Content-Length, then reads the exact body.
425 /// This avoids blocking on `read_to_string` when the server doesn't close
426 /// the connection immediately.
427 fn send_request(port: u16, path: &str) -> String {
428 let request = format!("GET {} HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n", path);
429 send_and_read(port, &request)
430 }
431
432 /// Core send-and-read implementation.
433 ///
434 /// The callback server uses a non-blocking accept loop (100ms poll interval),
435 /// so under heavy CPU load the server thread may not be scheduled immediately.
436 /// A small initial delay helps avoid connecting before the server is ready.
437 fn send_and_read(port: u16, request: &str) -> String {
438 // Give the server thread time to enter its accept loop.
439 // The server polls every 100ms, so 150ms ensures at least one accept cycle.
440 std::thread::sleep(Duration::from_millis(150));
441
442 let mut stream = TcpStream::connect(format!("127.0.0.1:{}", port)).unwrap();
443 stream
444 .set_read_timeout(Some(Duration::from_secs(5)))
445 .unwrap();
446 stream.write_all(request.as_bytes()).unwrap();
447 stream.flush().unwrap();
448
449 // Read byte-by-byte until we find \r\n\r\n (end of headers)
450 let mut raw = Vec::new();
451 let mut found_end = false;
452 loop {
453 let mut byte = [0u8; 1];
454 match stream.read(&mut byte) {
455 Ok(0) => break,
456 Ok(_) => {
457 raw.push(byte[0]);
458 if raw.len() >= 4 && &raw[raw.len() - 4..] == b"\r\n\r\n" {
459 found_end = true;
460 break;
461 }
462 }
463 Err(_) => break,
464 }
465 }
466
467 if !found_end {
468 return String::from_utf8_lossy(&raw).to_string();
469 }
470
471 let headers_str = String::from_utf8_lossy(&raw).to_string();
472
473 // Parse Content-Length from headers
474 let content_length: usize = headers_str
475 .lines()
476 .find_map(|line| {
477 let lower = line.to_lowercase();
478 if lower.starts_with("content-length:") {
479 lower
480 .strip_prefix("content-length:")
481 .and_then(|v| v.trim().parse().ok())
482 } else {
483 None
484 }
485 })
486 .unwrap_or(0);
487
488 // Read exactly content_length bytes for the body
489 let mut body_buf = vec![0u8; content_length];
490 if content_length > 0 {
491 let mut read_so_far = 0;
492 while read_so_far < content_length {
493 match stream.read(&mut body_buf[read_so_far..]) {
494 Ok(0) => break,
495 Ok(n) => read_so_far += n,
496 Err(_) => break,
497 }
498 }
499 }
500
501 let body_str = String::from_utf8_lossy(&body_buf);
502 format!("{}{}", headers_str, body_str)
503 }
504
505 /// Helper: send a raw HTTP request with a custom method and read the response.
506 fn send_raw_request(port: u16, request: &str) -> String {
507 send_and_read(port, request)
508 }
509
510 /// Helper: extract the HTTP body from a raw response (everything after \r\n\r\n).
511 fn extract_body(response: &str) -> &str {
512 response
513 .split("\r\n\r\n")
514 .nth(1)
515 .unwrap_or("")
516 }
517
518 /// Helper: extract the HTTP status line from a raw response.
519 fn extract_status(response: &str) -> &str {
520 response.lines().next().unwrap_or("")
521 }
522
523 #[test]
524 fn server_starts_on_random_port() {
525 let _lock = lock_server_tests();
526 let server = OAuthCallbackServer::start().unwrap();
527 assert!(server.port() > 0);
528 }
529
530 #[test]
531 fn server_returns_success_on_valid_callback() {
532 let _lock = lock_server_tests();
533 let server = OAuthCallbackServer::start().unwrap();
534 let port = server.port();
535
536 let response = send_request(port, "/?code=test_code_abc&state=test_state_xyz");
537 let body = extract_body(&response);
538
539 // Should return success HTML page
540 assert!(body.contains("Authorization Successful"));
541 assert!(body.contains("email account has been connected"));
542
543 // Should deliver result via channel
544 let result = server.wait_for_callback(Duration::from_secs(2));
545 let callback = result.unwrap().unwrap();
546 assert_eq!(callback.code, "test_code_abc");
547 assert_eq!(callback.state, "test_state_xyz");
548 }
549
550 #[test]
551 fn server_returns_error_on_oauth_error_callback() {
552 let _lock = lock_server_tests();
553 let server = OAuthCallbackServer::start().unwrap();
554 let port = server.port();
555
556 let response = send_request(
557 port,
558 "/?error=access_denied&error_description=User%20denied%20access",
559 );
560 let body = extract_body(&response);
561
562 // Should return error HTML page
563 assert!(body.contains("Authorization Failed"));
564 assert!(body.contains("User denied access"));
565
566 // Should deliver error via channel
567 let result = server.wait_for_callback(Duration::from_secs(2));
568 let err = result.unwrap().unwrap_err();
569 assert_eq!(err.error, "access_denied");
570 assert_eq!(err.error_description.as_deref(), Some("User denied access"));
571 }
572
573 #[test]
574 fn server_returns_error_without_description() {
575 let _lock = lock_server_tests();
576 let server = OAuthCallbackServer::start().unwrap();
577 let port = server.port();
578
579 let response = send_request(port, "/?error=server_error");
580 let body = extract_body(&response);
581
582 assert!(body.contains("Authorization Failed"));
583
584 let result = server.wait_for_callback(Duration::from_secs(2));
585 let err = result.unwrap().unwrap_err();
586 assert_eq!(err.error, "server_error");
587 assert!(err.error_description.is_none());
588 }
589
590 #[test]
591 fn server_rejects_non_get_methods() {
592 let _lock = lock_server_tests();
593 let server = OAuthCallbackServer::start().unwrap();
594 let port = server.port();
595
596 let response = send_raw_request(port, "POST / HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n");
597
598 assert!(extract_status(&response).contains("405"));
599 assert!(extract_body(&response).contains("Only GET is supported"));
600 }
601
602 #[test]
603 fn poll_returns_pending_initially() {
604 let _lock = lock_server_tests();
605 let server = OAuthCallbackServer::start().unwrap();
606
607 assert!(matches!(server.poll(), StoredCallback::Pending));
608 }
609
610 #[test]
611 fn poll_returns_success_after_callback() {
612 let _lock = lock_server_tests();
613 let server = OAuthCallbackServer::start().unwrap();
614 let port = server.port();
615
616 // Trigger the callback via the browser-redirect path.
617 send_request(port, "/?code=mycode&state=mystate");
618
619 // Small delay to let the server process.
620 std::thread::sleep(Duration::from_millis(200));
621
622 // The result is surfaced through poll(), not an HTTP endpoint.
623 match server.poll() {
624 StoredCallback::Success { code, state } => {
625 assert_eq!(code, "mycode");
626 assert_eq!(state, "mystate");
627 }
628 other => panic!("expected success, got {:?}", other),
629 }
630 }
631
632 #[test]
633 fn poll_returns_error_after_error_callback() {
634 let _lock = lock_server_tests();
635 let server = OAuthCallbackServer::start().unwrap();
636 let port = server.port();
637
638 // Trigger an error callback.
639 send_request(port, "/?error=invalid_grant&error_description=Expired");
640
641 std::thread::sleep(Duration::from_millis(200));
642
643 match server.poll() {
644 StoredCallback::Error { error, description } => {
645 assert_eq!(error, "invalid_grant");
646 assert_eq!(description.as_deref(), Some("Expired"));
647 }
648 other => panic!("expected error, got {:?}", other),
649 }
650 }
651
652 #[test]
653 fn callback_response_omits_cors_header() {
654 let _lock = lock_server_tests();
655 let server = OAuthCallbackServer::start().unwrap();
656 let port = server.port();
657
658 // The loopback server must not advertise a permissive CORS policy.
659 let response = send_request(port, "/?code=c&state=s");
660 assert!(!response.to_ascii_lowercase().contains("access-control-allow-origin"));
661 }
662
663 #[test]
664 fn server_timeout_returns_error() {
665 let _lock = lock_server_tests();
666 let server = OAuthCallbackServer::start().unwrap();
667
668 // Don't send any request, just wait with a very short timeout
669 let result = server.wait_for_callback(Duration::from_millis(50));
670 assert!(result.is_err());
671 assert!(result.unwrap_err().contains("Timeout"));
672 }
673
674 #[test]
675 fn server_ignores_request_with_no_code_or_error() {
676 let _lock = lock_server_tests();
677 let server = OAuthCallbackServer::start().unwrap();
678 let port = server.port();
679
680 // Send a request to root with no query params.
681 // The handler returns None early (no code/error), so no HTTP response is sent.
682 // We just connect and send the request without expecting a response.
683 let mut stream = TcpStream::connect(format!("127.0.0.1:{}", port)).unwrap();
684 stream.set_write_timeout(Some(Duration::from_secs(2))).unwrap();
685 let request = "GET / HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n";
686 stream.write_all(request.as_bytes()).unwrap();
687 stream.flush().unwrap();
688 drop(stream);
689
690 // Give the server a moment to process
691 std::thread::sleep(Duration::from_millis(100));
692
693 // The channel should still be empty (no callback delivered)
694 let result = server.wait_for_callback(Duration::from_millis(100));
695 assert!(result.is_err());
696 }
697
698 #[test]
699 fn server_handles_code_without_state() {
700 let _lock = lock_server_tests();
701 let server = OAuthCallbackServer::start().unwrap();
702 let port = server.port();
703
704 // Code present but state missing -- should not produce a callback
705 send_request(port, "/?code=only_code");
706
707 let result = server.wait_for_callback(Duration::from_millis(100));
708 assert!(result.is_err());
709 }
710
711 #[test]
712 fn server_handles_state_without_code() {
713 let _lock = lock_server_tests();
714 let server = OAuthCallbackServer::start().unwrap();
715 let port = server.port();
716
717 // State present but code missing -- should not produce a callback
718 send_request(port, "/?state=only_state");
719
720 let result = server.wait_for_callback(Duration::from_millis(100));
721 assert!(result.is_err());
722 }
723 }
724