@@ -2816,6 +2816,17 @@ mod tests {
28162816 use super :: * ;
28172817 use aws_smithy_http_client:: test_util:: { CaptureRequestReceiver , capture_request} ;
28182818 use std:: collections:: HashMap ;
2819+ use std:: io:: { Read , Write } ;
2820+ use std:: net:: { TcpListener , TcpStream } ;
2821+ use std:: sync:: mpsc;
2822+ use std:: thread;
2823+
2824+ #[ derive( Debug ) ]
2825+ struct CapturedXmlRequest {
2826+ method : String ,
2827+ target : String ,
2828+ headers : Vec < ( String , String ) > ,
2829+ }
28192830
28202831 fn test_s3_client (
28212832 response : Option < http:: Response < SdkBody > > ,
@@ -2873,6 +2884,69 @@ mod tests {
28732884 ( client, request_receiver)
28742885 }
28752886
2887+ fn read_xml_request ( stream : & mut TcpStream ) -> CapturedXmlRequest {
2888+ let mut buffer = Vec :: new ( ) ;
2889+ let mut chunk = [ 0_u8 ; 1024 ] ;
2890+ let header_end = loop {
2891+ let read = stream. read ( & mut chunk) . expect ( "read HTTP request" ) ;
2892+ assert ! ( read > 0 , "client closed connection before headers" ) ;
2893+ buffer. extend_from_slice ( & chunk[ ..read] ) ;
2894+
2895+ if let Some ( position) = buffer. windows ( 4 ) . position ( |window| window == b"\r \n \r \n " ) {
2896+ break position + 4 ;
2897+ }
2898+ } ;
2899+
2900+ let headers_text = String :: from_utf8_lossy ( & buffer[ ..header_end] ) . into_owned ( ) ;
2901+ let mut lines = headers_text. lines ( ) ;
2902+ let request_line = lines. next ( ) . expect ( "request line" ) ;
2903+ let mut parts = request_line. split_whitespace ( ) ;
2904+ let method = parts. next ( ) . expect ( "request method" ) . to_string ( ) ;
2905+ let target = parts. next ( ) . expect ( "request target" ) . to_string ( ) ;
2906+ let headers = lines
2907+ . filter_map ( |line| {
2908+ let ( name, value) = line. split_once ( ':' ) ?;
2909+ Some ( ( name. to_ascii_lowercase ( ) , value. trim ( ) . to_string ( ) ) )
2910+ } )
2911+ . collect ( ) ;
2912+
2913+ CapturedXmlRequest {
2914+ method,
2915+ target,
2916+ headers,
2917+ }
2918+ }
2919+
2920+ fn start_xml_test_server ( ) -> (
2921+ String ,
2922+ mpsc:: Receiver < CapturedXmlRequest > ,
2923+ thread:: JoinHandle < ( ) > ,
2924+ ) {
2925+ let listener = TcpListener :: bind ( "127.0.0.1:0" ) . expect ( "bind test server" ) ;
2926+ let endpoint = format ! ( "http://{}" , listener. local_addr( ) . expect( "local addr" ) ) ;
2927+ let ( sender, receiver) = mpsc:: channel ( ) ;
2928+
2929+ let handle = thread:: spawn ( move || {
2930+ let ( mut stream, _) = listener. accept ( ) . expect ( "accept request" ) ;
2931+ let request = read_xml_request ( & mut stream) ;
2932+ sender. send ( request) . expect ( "send captured request" ) ;
2933+
2934+ let response = "HTTP/1.1 200 OK\r \n content-length: 2\r \n connection: close\r \n \r \n ok" ;
2935+ stream
2936+ . write_all ( response. as_bytes ( ) )
2937+ . expect ( "write HTTP response" ) ;
2938+ } ) ;
2939+
2940+ ( endpoint, receiver, handle)
2941+ }
2942+
2943+ fn header_value < ' a > ( headers : & ' a [ ( String , String ) ] , name : & str ) -> Option < & ' a str > {
2944+ headers
2945+ . iter ( )
2946+ . find ( |( header_name, _) | header_name. eq_ignore_ascii_case ( name) )
2947+ . map ( |( _, value) | value. as_str ( ) )
2948+ }
2949+
28762950 #[ test]
28772951 fn test_object_info_creation ( ) {
28782952 let info = ObjectInfo :: file ( "test.txt" , 1024 ) ;
@@ -3595,6 +3669,49 @@ mod tests {
35953669 assert ! ( !url. contains( "x-amz-bucket-encrypt-enabled" ) ) ;
35963670 }
35973671
3672+ #[ tokio:: test]
3673+ async fn custom_headers_are_added_to_xml_requests_before_signing ( ) {
3674+ let ( endpoint, request_receiver, server_handle) = start_xml_test_server ( ) ;
3675+ let ( client, _sdk_request_receiver) = test_s3_client_with_endpoint_and_headers (
3676+ & endpoint,
3677+ None ,
3678+ vec ! [ RequestHeader {
3679+ name: "x-amz-bucket-encrypt-enabled" . to_string( ) ,
3680+ value: "1" . to_string( ) ,
3681+ } ] ,
3682+ ) ;
3683+ let url = client
3684+ . replication_url ( "bucket" )
3685+ . expect ( "replication URL should build" ) ;
3686+
3687+ let response = client
3688+ . xml_request (
3689+ Method :: PUT ,
3690+ url,
3691+ Some ( "application/xml" ) ,
3692+ Some ( b"<xml/>" . to_vec ( ) ) ,
3693+ )
3694+ . await
3695+ . expect ( "xml request should succeed" ) ;
3696+
3697+ assert_eq ! ( response, "ok" ) ;
3698+ let request = request_receiver
3699+ . recv ( )
3700+ . expect ( "server should capture XML request" ) ;
3701+ assert_eq ! ( request. method, "PUT" ) ;
3702+ assert_eq ! ( request. target, "/bucket?replication=" ) ;
3703+ assert_eq ! (
3704+ header_value( & request. headers, "x-amz-bucket-encrypt-enabled" ) ,
3705+ Some ( "1" )
3706+ ) ;
3707+ assert ! (
3708+ header_value( & request. headers, "authorization" )
3709+ . expect( "authorization header" )
3710+ . contains( "x-amz-bucket-encrypt-enabled" )
3711+ ) ;
3712+ server_handle. join ( ) . expect ( "server thread should finish" ) ;
3713+ }
3714+
35983715 #[ tokio:: test]
35993716 async fn delete_object_without_force_delete_omits_rustfs_header ( ) {
36003717 let ( client, request_receiver) = test_s3_client ( None ) ;
0 commit comments