@@ -531,14 +531,74 @@ fn parse_query_params(query: &str) -> HashMap<String, String> {
531531#[ cfg( test) ]
532532mod 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