Skip to content

Commit 3402ff5

Browse files
committed
Expand authorize.rs test coverage across callback paths
1 parent 982d733 commit 3402ff5

1 file changed

Lines changed: 145 additions & 1 deletion

File tree

src/authorize.rs

Lines changed: 145 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -531,14 +531,74 @@ fn parse_query_params(query: &str) -> HashMap<String, String> {
531531
#[cfg(test)]
532532
mod tests {
533533
use super::*;
534-
use std::net::TcpListener as StdTcpListener;
534+
use std::io::{Read, Write};
535+
use std::net::{TcpListener as StdTcpListener, TcpStream};
536+
use std::sync::mpsc;
537+
use std::thread;
538+
use std::time::Duration as StdDuration;
539+
use tokio::runtime::Runtime;
535540
use tokio::time::{timeout, Duration};
536541

537542
fn reserve_ephemeral_port() -> u16 {
538543
let listener = StdTcpListener::bind("127.0.0.1:0").expect("failed to bind ephemeral port");
539544
listener.local_addr().expect("failed to get local addr").port()
540545
}
541546

547+
fn spawn_callback_server(
548+
port: u16,
549+
auth_code: Arc<Mutex<Option<String>>>,
550+
) -> mpsc::Receiver<Result<String, String>> {
551+
let (tx, rx) = mpsc::channel();
552+
thread::spawn(move || {
553+
let runtime = Runtime::new().expect("failed to create tokio runtime");
554+
let result = runtime
555+
.block_on(start_callback_server(port, auth_code))
556+
.map_err(|e| e.to_string());
557+
tx.send(result).expect("failed to send callback result");
558+
});
559+
560+
rx
561+
}
562+
563+
fn send_http_get(port: u16, path: &str) -> (u16, String) {
564+
let mut stream = None;
565+
566+
for _ in 0..50 {
567+
match TcpStream::connect(("127.0.0.1", port)) {
568+
Ok(s) => {
569+
stream = Some(s);
570+
break;
571+
}
572+
Err(_) => thread::sleep(StdDuration::from_millis(20)),
573+
}
574+
}
575+
576+
let mut stream = stream.expect("failed to connect to callback server");
577+
let request = format!("GET {} HTTP/1.0\r\n\r\n", path);
578+
579+
stream
580+
.write_all(request.as_bytes())
581+
.expect("failed to write request");
582+
583+
let mut raw_response = String::new();
584+
stream
585+
.read_to_string(&mut raw_response)
586+
.expect("failed to read response");
587+
588+
let mut sections = raw_response.splitn(2, "\r\n\r\n");
589+
let headers = sections.next().expect("response headers missing");
590+
let body = sections.next().unwrap_or_default().to_string();
591+
let status_line = headers.lines().next().expect("status line missing");
592+
let status = status_line
593+
.split_whitespace()
594+
.nth(1)
595+
.expect("status code missing")
596+
.parse::<u16>()
597+
.expect("invalid status code");
598+
599+
(status, body)
600+
}
601+
542602
#[test]
543603
fn parse_query_params_decodes_values() {
544604
let params = parse_query_params("code=a%20b&error_description=needs%2Blogin");
@@ -547,6 +607,41 @@ mod tests {
547607
assert_eq!(params.get("error_description"), Some(&"needs+login".to_string()));
548608
}
549609

610+
#[test]
611+
fn parse_query_params_ignores_malformed_pairs() {
612+
let params = parse_query_params("valid=ok&invalid&also_invalid=");
613+
614+
assert_eq!(params.get("valid"), Some(&"ok".to_string()));
615+
assert_eq!(params.get("invalid"), None);
616+
assert_eq!(params.get("also_invalid"), Some(&"".to_string()));
617+
}
618+
619+
#[test]
620+
fn port_is_available_reflects_current_port_usage() {
621+
let listener = StdTcpListener::bind("127.0.0.1:0").expect("failed to bind ephemeral port");
622+
let port = listener
623+
.local_addr()
624+
.expect("failed to get listener addr")
625+
.port();
626+
627+
assert!(!port_is_available(port));
628+
drop(listener);
629+
assert!(port_is_available(port));
630+
}
631+
632+
#[test]
633+
fn find_available_port_skips_ports_that_are_in_use() {
634+
let listener = StdTcpListener::bind("127.0.0.1:0").expect("failed to bind ephemeral port");
635+
let occupied_port = listener
636+
.local_addr()
637+
.expect("failed to get listener addr")
638+
.port();
639+
640+
let found_port = find_available_port(occupied_port).expect("should find an available port");
641+
642+
assert_ne!(found_port, occupied_port);
643+
}
644+
550645
#[tokio::test]
551646
async fn start_callback_server_returns_without_waiting_for_second_connection() {
552647
let port = reserve_ephemeral_port();
@@ -562,4 +657,53 @@ mod tests {
562657

563658
assert_eq!(returned_code, "test-code");
564659
}
660+
661+
#[test]
662+
fn start_callback_server_returns_bind_error_if_port_is_occupied() {
663+
let listener = StdTcpListener::bind("127.0.0.1:0").expect("failed to bind ephemeral port");
664+
let occupied_port = listener
665+
.local_addr()
666+
.expect("failed to get listener addr")
667+
.port();
668+
669+
let runtime = Runtime::new().expect("failed to create runtime");
670+
let result = runtime.block_on(start_callback_server(
671+
occupied_port,
672+
Arc::new(Mutex::new(None::<String>)),
673+
));
674+
675+
assert!(result.is_err());
676+
let error = result.err().expect("expected bind error").to_string();
677+
assert!(error.contains("Failed to bind"));
678+
}
679+
680+
#[test]
681+
fn callback_server_serves_waiting_error_and_success_pages_then_returns_code() {
682+
let port = reserve_ephemeral_port();
683+
let auth_code = Arc::new(Mutex::new(None::<String>));
684+
let result_rx = spawn_callback_server(port, auth_code);
685+
686+
let (waiting_status, waiting_body) = send_http_get(port, "/");
687+
assert_eq!(waiting_status, 200);
688+
assert!(waiting_body.contains("Waiting for Authorization"));
689+
690+
let (error_status, error_body) = send_http_get(
691+
port,
692+
"/?error=access_denied&error_description=user%20cancelled",
693+
);
694+
assert_eq!(error_status, 400);
695+
assert!(error_body.contains("Authorization Failed"));
696+
assert!(error_body.contains("access_denied"));
697+
698+
let (success_status, success_body) = send_http_get(port, "/?code=abc123");
699+
assert_eq!(success_status, 200);
700+
assert!(success_body.contains("Successfully Signed In"));
701+
702+
let returned_code = result_rx
703+
.recv_timeout(StdDuration::from_secs(2))
704+
.expect("callback server should return in time")
705+
.expect("callback server should return code");
706+
707+
assert_eq!(returned_code, "abc123");
708+
}
565709
}

0 commit comments

Comments
 (0)