@@ -352,8 +352,8 @@ where
352352 "Fail to proxy: {e}, downstream session type: {}" ,
353353 downstream_session. session_type( )
354354 ) ;
355- downstream_session
356- . respond_error ( 400 )
355+ self . inner
356+ . request_error_filter ( & mut downstream_session , & e )
357357 . await
358358 . unwrap_or_else ( |e| {
359359 error ! ( "failed to send error response to downstream: {e}" ) ;
@@ -1891,6 +1891,7 @@ mod tests {
18911891 read_buf : Vec < u8 > ,
18921892 read_pos : usize ,
18931893 write_buf : Arc < Mutex < Vec < u8 > > > ,
1894+ shutdown : Arc < AtomicBool > ,
18941895 }
18951896
18961897 impl StaticVirtualSocket {
@@ -1899,6 +1900,7 @@ mod tests {
18991900 read_buf : read_buf. to_vec ( ) ,
19001901 read_pos : 0 ,
19011902 write_buf,
1903+ shutdown : Arc :: new ( AtomicBool :: new ( false ) ) ,
19021904 }
19031905 }
19041906 }
@@ -1934,6 +1936,7 @@ mod tests {
19341936 }
19351937
19361938 fn poll_shutdown ( self : Pin < & mut Self > , _cx : & mut Context < ' _ > ) -> Poll < std:: io:: Result < ( ) > > {
1939+ self . shutdown . store ( true , Ordering :: Relaxed ) ;
19371940 Poll :: Ready ( Ok ( ( ) ) )
19381941 }
19391942 }
@@ -1977,6 +1980,229 @@ mod tests {
19771980 }
19781981 }
19791982
1983+ fn unread_request_session (
1984+ request : & [ u8 ] ,
1985+ ) -> ( HttpSession , Arc < Mutex < Vec < u8 > > > , Arc < AtomicBool > ) {
1986+ let written = Arc :: new ( Mutex :: new ( Vec :: new ( ) ) ) ;
1987+ let socket = StaticVirtualSocket :: new ( request, written. clone ( ) ) ;
1988+ let shutdown = socket. shutdown . clone ( ) ;
1989+ let stream = L4Stream :: from ( VirtualSocketStream :: new ( Box :: new ( socket) ) ) ;
1990+ ( HttpSession :: new_http1 ( Box :: new ( stream) ) , written, shutdown)
1991+ }
1992+
1993+ const MALFORMED_REQUEST : & [ u8 ] = b"GET / HTTP/1.1\r \n Bad Header: value\r \n \r \n " ;
1994+ const INVALID_REQUEST : & [ u8 ] =
1995+ b"GET / HTTP/1.1\r \n Host: example.com\r \n Content-Length: 1\r \n Content-Length: 2\r \n \r \n " ;
1996+ const CUSTOM_ERROR_BODY : & [ u8 ] = b"bad request\n " ;
1997+
1998+ enum RequestErrorAction {
1999+ Respond ,
2000+ Close ,
2001+ Fail ,
2002+ RespondThenFail ,
2003+ }
2004+
2005+ struct RequestErrorProxy {
2006+ calls : AtomicUsize ,
2007+ action : RequestErrorAction ,
2008+ }
2009+
2010+ #[ async_trait]
2011+ impl ProxyHttp for RequestErrorProxy {
2012+ type CTX = ( ) ;
2013+
2014+ fn new_ctx ( & self ) -> Self :: CTX {
2015+ panic ! ( "rejected requests must not create a proxy context" ) ;
2016+ }
2017+
2018+ async fn upstream_peer (
2019+ & self ,
2020+ _session : & mut Session ,
2021+ _ctx : & mut Self :: CTX ,
2022+ ) -> Result < Box < HttpPeer > > {
2023+ unreachable ! ( "rejected requests must not reach upstream selection" ) ;
2024+ }
2025+
2026+ async fn request_filter (
2027+ & self ,
2028+ _session : & mut Session ,
2029+ _ctx : & mut Self :: CTX ,
2030+ ) -> Result < bool > {
2031+ unreachable ! ( "rejected requests must not reach request filters" ) ;
2032+ }
2033+
2034+ async fn logging ( & self , _session : & mut Session , _e : Option < & Error > , _ctx : & mut Self :: CTX ) {
2035+ unreachable ! ( "rejected requests must not reach normal logging" ) ;
2036+ }
2037+
2038+ async fn request_error_filter ( & self , session : & mut HttpSession , e : & Error ) -> Result < ( ) > {
2039+ assert_eq ! ( e. etype( ) , & InvalidHTTPHeader ) ;
2040+ assert_eq ! ( e. esource( ) , & ErrorSource :: Downstream ) ;
2041+ self . calls . fetch_add ( 1 , Ordering :: Relaxed ) ;
2042+ // These accessors are usable even if parsing failed before a header was created.
2043+ let _ = session. client_addr ( ) ;
2044+ let _ = session. digest ( ) ;
2045+ if matches ! ( self . action, RequestErrorAction :: Close ) {
2046+ return Ok ( ( ) ) ;
2047+ }
2048+ if matches ! ( self . action, RequestErrorAction :: Fail ) {
2049+ return Err ( Error :: new ( WriteError ) ) ;
2050+ }
2051+ let mut response = ResponseHeader :: build ( 422 , Some ( 2 ) ) ?;
2052+ response. set_content_length ( CUSTOM_ERROR_BODY . len ( ) ) ?;
2053+ response. insert_header ( header:: CONTENT_TYPE , "text/plain" ) ?;
2054+ session
2055+ . write_error_response ( response, Bytes :: from_static ( CUSTOM_ERROR_BODY ) )
2056+ . await ?;
2057+ if matches ! ( self . action, RequestErrorAction :: RespondThenFail ) {
2058+ return Err ( Error :: new ( WriteError ) ) ;
2059+ }
2060+ Ok ( ( ) )
2061+ }
2062+ }
2063+
2064+ fn request_error_proxy ( action : RequestErrorAction ) -> Arc < HttpProxy < RequestErrorProxy > > {
2065+ Arc :: new ( HttpProxy :: new (
2066+ RequestErrorProxy {
2067+ calls : AtomicUsize :: new ( 0 ) ,
2068+ action,
2069+ } ,
2070+ Arc :: new ( ServerConf :: default ( ) ) ,
2071+ ) )
2072+ }
2073+
2074+ fn assert_rejection_response (
2075+ written : & Mutex < Vec < u8 > > ,
2076+ status : u16 ,
2077+ body : & [ u8 ] ,
2078+ content_length : usize ,
2079+ ) {
2080+ let written = written. lock ( ) . unwrap ( ) ;
2081+ let mut headers = [ httparse:: EMPTY_HEADER ; 16 ] ;
2082+ let mut response = httparse:: Response :: new ( & mut headers) ;
2083+ let httparse:: Status :: Complete ( header_len) = response. parse ( & written) . unwrap ( ) else {
2084+ panic ! ( "incomplete error response" ) ;
2085+ } ;
2086+ assert_eq ! ( response. code, Some ( status) ) ;
2087+ let header = |name : & str | {
2088+ response
2089+ . headers
2090+ . iter ( )
2091+ . find ( |h| h. name . eq_ignore_ascii_case ( name) )
2092+ . map ( |h| h. value )
2093+ } ;
2094+ assert_eq ! ( header( "connection" ) , Some ( b"close" . as_slice( ) ) ) ;
2095+ assert_eq ! (
2096+ header( "content-length" ) ,
2097+ Some ( content_length. to_string( ) . as_bytes( ) )
2098+ ) ;
2099+ if status == 422 {
2100+ assert_eq ! ( header( "content-type" ) , Some ( b"text/plain" . as_slice( ) ) ) ;
2101+ assert ! ( header( "server" ) . is_none( ) ) ;
2102+ } else {
2103+ assert ! ( header( "server" ) . is_some( ) ) ;
2104+ }
2105+ // Also rules out an appended fallback response or a response to pipelined input.
2106+ assert_eq ! ( & written[ header_len..] , body) ;
2107+ }
2108+
2109+ #[ tokio:: test]
2110+ async fn request_error_filter_defaults_to_400 ( ) {
2111+ let proxy = HttpProxy :: new ( DefaultRetryProxy , Arc :: new ( ServerConf :: default ( ) ) ) ;
2112+ for request in [ MALFORMED_REQUEST , INVALID_REQUEST ] {
2113+ let ( session, written, shutdown) = unread_request_session ( request) ;
2114+ assert ! ( proxy. handle_new_request( Box :: new( session) ) . await . is_none( ) ) ;
2115+ assert ! ( shutdown. load( Ordering :: Relaxed ) ) ;
2116+ assert_rejection_response ( & written, 400 , b"" , 0 ) ;
2117+ }
2118+ }
2119+
2120+ #[ tokio:: test]
2121+ async fn request_error_filter_customizes_rejected_requests_before_context_creation ( ) {
2122+ for request in [ MALFORMED_REQUEST , INVALID_REQUEST ] {
2123+ let proxy = request_error_proxy ( RequestErrorAction :: Respond ) ;
2124+ let ( session, written, closed) = unread_request_session ( request) ;
2125+ let ( _shutdown_tx, shutdown) = tokio:: sync:: watch:: channel ( false ) ;
2126+ assert ! ( proxy. process_new_http( session, & shutdown) . await . is_none( ) ) ;
2127+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 1 ) ;
2128+ assert ! ( closed. load( Ordering :: Relaxed ) ) ;
2129+ assert_rejection_response ( & written, 422 , CUSTOM_ERROR_BODY , CUSTOM_ERROR_BODY . len ( ) ) ;
2130+ }
2131+ }
2132+
2133+ #[ tokio:: test]
2134+ async fn request_error_filter_preserves_head_body_suppression ( ) {
2135+ let request = b"HEAD / HTTP/1.1\r \n Host: example.com\r \n Content-Length: 1\r \n Content-Length: 2\r \n \r \n " ;
2136+ let proxy = request_error_proxy ( RequestErrorAction :: Respond ) ;
2137+ let ( session, written, shutdown) = unread_request_session ( request) ;
2138+ assert ! ( proxy. handle_new_request( Box :: new( session) ) . await . is_none( ) ) ;
2139+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 1 ) ;
2140+ assert ! ( shutdown. load( Ordering :: Relaxed ) ) ;
2141+ assert_rejection_response ( & written, 422 , b"" , CUSTOM_ERROR_BODY . len ( ) ) ;
2142+ }
2143+
2144+ #[ tokio:: test]
2145+ async fn request_error_filter_closes_without_a_fallback_response ( ) {
2146+ for action in [ RequestErrorAction :: Close , RequestErrorAction :: Fail ] {
2147+ let proxy = request_error_proxy ( action) ;
2148+ let ( session, written, shutdown) = unread_request_session ( MALFORMED_REQUEST ) ;
2149+ assert ! ( proxy. handle_new_request( Box :: new( session) ) . await . is_none( ) ) ;
2150+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 1 ) ;
2151+ assert ! ( shutdown. load( Ordering :: Relaxed ) ) ;
2152+ assert ! ( written. lock( ) . unwrap( ) . is_empty( ) ) ;
2153+ }
2154+ }
2155+
2156+ #[ tokio:: test]
2157+ async fn request_error_filter_does_not_append_a_response_after_callback_failure ( ) {
2158+ let proxy = request_error_proxy ( RequestErrorAction :: RespondThenFail ) ;
2159+ let ( session, written, shutdown) = unread_request_session ( MALFORMED_REQUEST ) ;
2160+ assert ! ( proxy. handle_new_request( Box :: new( session) ) . await . is_none( ) ) ;
2161+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 1 ) ;
2162+ assert ! ( shutdown. load( Ordering :: Relaxed ) ) ;
2163+ assert_rejection_response ( & written, 422 , CUSTOM_ERROR_BODY , CUSTOM_ERROR_BODY . len ( ) ) ;
2164+ }
2165+
2166+ #[ tokio:: test]
2167+ async fn request_error_filter_does_not_process_pipelined_input_after_rejection ( ) {
2168+ let mut request = MALFORMED_REQUEST . to_vec ( ) ;
2169+ request. extend_from_slice ( b"GET / HTTP/1.1\r \n Host: example.com\r \n \r \n " ) ;
2170+ let proxy = request_error_proxy ( RequestErrorAction :: Respond ) ;
2171+ let ( session, written, closed) = unread_request_session ( & request) ;
2172+ let ( _shutdown_tx, shutdown) = tokio:: sync:: watch:: channel ( false ) ;
2173+ assert ! ( proxy. process_new_http( session, & shutdown) . await . is_none( ) ) ;
2174+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 1 ) ;
2175+ assert ! ( closed. load( Ordering :: Relaxed ) ) ;
2176+ assert_rejection_response ( & written, 422 , CUSTOM_ERROR_BODY , CUSTOM_ERROR_BODY . len ( ) ) ;
2177+ }
2178+
2179+ #[ tokio:: test]
2180+ async fn request_error_filter_skips_valid_requests_and_connection_errors ( ) {
2181+ for request in [
2182+ b"GET / HTTP/1.1\r \n Host: example.com\r \n \r \n " . as_slice ( ) ,
2183+ b"" ,
2184+ b"GET / HTTP/1.1\r \n Host:" ,
2185+ ] {
2186+ let proxy = request_error_proxy ( RequestErrorAction :: Respond ) ;
2187+ let ( session, written, shutdown) = unread_request_session ( request) ;
2188+ let result = proxy. handle_new_request ( Box :: new ( session) ) . await ;
2189+ assert_eq ! ( result. is_some( ) , request. ends_with( b"\r \n \r \n " ) ) ;
2190+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 0 ) ;
2191+ assert ! ( written. lock( ) . unwrap( ) . is_empty( ) ) ;
2192+ if request. ends_with ( b"Host:" ) {
2193+ assert ! ( shutdown. load( Ordering :: Relaxed ) ) ;
2194+ }
2195+ }
2196+ }
2197+
2198+ #[ tokio:: test]
2199+ async fn request_error_filter_skips_shutdown ( ) {
2200+ let proxy = request_error_proxy ( RequestErrorAction :: Respond ) ;
2201+ proxy. http_cleanup ( ) . await ;
2202+ assert ! ( proxy. handle_new_request( pending_session( ) ) . await . is_none( ) ) ;
2203+ assert_eq ! ( proxy. inner. calls. load( Ordering :: Relaxed ) , 0 ) ;
2204+ }
2205+
19802206 fn default_policy_would_retry_for_session (
19812207 session : & mut Session ,
19822208 retry : RetryType ,
0 commit comments