diff --git a/src/client/conn/http1.rs b/src/client/conn/http1.rs index bdb03171d4..88063a6d71 100644 --- a/src/client/conn/http1.rs +++ b/src/client/conn/http1.rs @@ -131,6 +131,7 @@ pub struct Builder { h1_title_case_headers: bool, h1_preserve_header_case: bool, h1_max_headers: Option, + h1_max_header_size: Option, #[cfg(feature = "ffi")] h1_preserve_header_order: bool, h1_read_buf_exact_size: Option, @@ -352,6 +353,7 @@ impl Builder { h1_title_case_headers: false, h1_preserve_header_case: false, h1_max_headers: None, + h1_max_header_size: None, #[cfg(feature = "ffi")] h1_preserve_header_order: false, h1_max_buf_size: None, @@ -500,6 +502,17 @@ impl Builder { self } + /// Set the maximum size of response headers (including the status line) in bytes. + /// + /// If the server sends headers exceeding this limit, the error "message head is too large" + /// is returned. + /// + /// Default is `None`. + pub fn max_header_size(&mut self, val: usize) -> &mut Self { + self.h1_max_header_size = Some(val); + self + } + /// Set whether to support preserving original header order. /// /// Currently, this will record the order in which headers are received, and store this @@ -587,6 +600,9 @@ impl Builder { if let Some(max_headers) = opts.h1_max_headers { conn.set_http1_max_headers(max_headers); } + if let Some(max_header_size) = opts.h1_max_header_size { + conn.set_http1_max_header_size(max_header_size); + } #[cfg(feature = "ffi")] if opts.h1_preserve_header_order { conn.set_preserve_header_order(); diff --git a/src/proto/h1/conn.rs b/src/proto/h1/conn.rs index 593d3d9eda..b26eefd2a7 100644 --- a/src/proto/h1/conn.rs +++ b/src/proto/h1/conn.rs @@ -58,6 +58,7 @@ where method: None, h1_parser_config: ParserConfig::default(), h1_max_headers: None, + h1_max_header_size: None, #[cfg(feature = "server")] h1_header_read_timeout: None, #[cfg(feature = "server")] @@ -141,6 +142,10 @@ where self.state.h1_max_headers = Some(val); } + pub(crate) fn set_http1_max_header_size(&mut self, val: usize) { + self.state.h1_max_header_size = Some(val); + } + #[cfg(feature = "server")] pub(crate) fn set_http1_header_read_timeout(&mut self, val: Duration) { self.state.h1_header_read_timeout = Some(val); @@ -241,6 +246,7 @@ where req_method: &mut self.state.method, h1_parser_config: self.state.h1_parser_config.clone(), h1_max_headers: self.state.h1_max_headers, + h1_max_header_size: self.state.h1_max_header_size, preserve_header_case: self.state.preserve_header_case, #[cfg(feature = "ffi")] preserve_header_order: self.state.preserve_header_order, @@ -309,7 +315,7 @@ where self.try_keep_alive(cx); } } else if msg.expect_continue && msg.head.version.gt(&Version::HTTP_10) { - let h1_max_header_size = None; // TODO: remove this when we land h1_max_header_size support + let h1_max_header_size = self.state.h1_max_header_size; self.state.reading = Reading::Continue(Decoder::new( msg.decode, self.state.h1_max_headers, @@ -317,7 +323,7 @@ where )); wants = wants.add(Wants::EXPECT); } else { - let h1_max_header_size = None; // TODO: remove this when we land h1_max_header_size support + let h1_max_header_size = self.state.h1_max_header_size; self.state.reading = Reading::Body(Decoder::new( msg.decode, self.state.h1_max_headers, @@ -933,6 +939,7 @@ struct State { method: Option, h1_parser_config: ParserConfig, h1_max_headers: Option, + h1_max_header_size: Option, #[cfg(feature = "server")] h1_header_read_timeout: Option, #[cfg(feature = "server")] diff --git a/src/proto/h1/io.rs b/src/proto/h1/io.rs index b49e48e5c8..0dd31d3a22 100644 --- a/src/proto/h1/io.rs +++ b/src/proto/h1/io.rs @@ -188,6 +188,7 @@ where req_method: parse_ctx.req_method, h1_parser_config: parse_ctx.h1_parser_config.clone(), h1_max_headers: parse_ctx.h1_max_headers, + h1_max_header_size: parse_ctx.h1_max_header_size, preserve_header_case: parse_ctx.preserve_header_case, #[cfg(feature = "ffi")] preserve_header_order: parse_ctx.preserve_header_order, @@ -200,8 +201,14 @@ where self.partial_len = None; return Poll::Ready(Ok(msg)); } else { - let max = self.read_buf_strategy.max(); let curr_len = self.read_buf.len(); + if let Some(max_header_size) = parse_ctx.h1_max_header_size { + if curr_len >= max_header_size { + debug!("max_header_size ({}) reached, closing", max_header_size); + return Poll::Ready(Err(crate::Error::new_too_large())); + } + } + let max = self.read_buf_strategy.max(); if curr_len >= max { debug!("max_buf_size ({}) reached, closing", max); return Poll::Ready(Err(crate::Error::new_too_large())); @@ -702,6 +709,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, diff --git a/src/proto/h1/mod.rs b/src/proto/h1/mod.rs index a17dbae83c..344d5a8a4f 100644 --- a/src/proto/h1/mod.rs +++ b/src/proto/h1/mod.rs @@ -73,6 +73,7 @@ pub(crate) struct ParseContext<'ctx> { req_method: &'ctx mut Option, h1_parser_config: ParserConfig, h1_max_headers: Option, + h1_max_header_size: Option, preserve_header_case: bool, #[cfg(feature = "ffi")] preserve_header_order: bool, diff --git a/src/proto/h1/role.rs b/src/proto/h1/role.rs index d083d2a912..8aefcaf2d9 100644 --- a/src/proto/h1/role.rs +++ b/src/proto/h1/role.rs @@ -180,6 +180,11 @@ impl Http1Transaction for Server { ) { Ok(httparse::Status::Complete(parsed_len)) => { trace!("Request.parse Complete({})", parsed_len); + if let Some(max_header_size) = ctx.h1_max_header_size { + if parsed_len > max_header_size { + return Err(Parse::TooLarge); + } + } len = parsed_len; let uri = req.path.expect("httparse completed"); if uri.len() > MAX_URI_LEN { @@ -1052,6 +1057,11 @@ impl Http1Transaction for Client { ) { Ok(httparse::Status::Complete(len)) => { trace!("Response.parse Complete({})", len); + if let Some(max_header_size) = ctx.h1_max_header_size { + if len > max_header_size { + return Err(Parse::TooLarge); + } + } let status = StatusCode::from_u16(res.code.expect("httparse completed"))?; let reason = { @@ -1692,6 +1702,7 @@ mod tests { req_method: &mut method, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1720,6 +1731,7 @@ mod tests { req_method: &mut Some(crate::Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1744,6 +1756,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1765,6 +1778,7 @@ mod tests { req_method: &mut Some(crate::Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1788,6 +1802,7 @@ mod tests { req_method: &mut Some(crate::Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1815,6 +1830,7 @@ mod tests { req_method: &mut Some(crate::Method::GET), h1_parser_config, h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1839,6 +1855,7 @@ mod tests { req_method: &mut Some(crate::Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1867,6 +1884,7 @@ mod tests { req_method: &mut method, h1_parser_config, h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1894,6 +1912,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1914,6 +1933,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: true, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1953,6 +1973,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -1974,6 +1995,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2214,6 +2236,7 @@ mod tests { req_method: &mut Some(Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2235,6 +2258,7 @@ mod tests { req_method: &mut Some(m), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2256,6 +2280,7 @@ mod tests { req_method: &mut Some(Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2826,6 +2851,7 @@ mod tests { req_method: &mut Some(Method::GET), h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2870,6 +2896,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: max_headers, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2894,6 +2921,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: max_headers, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -2970,6 +2998,70 @@ mod tests { parse(Some(200), 210, false); } + #[test] + fn test_h1_max_header_size() { + let _ = pretty_env_logger::try_init(); + + let req_str = "GET / HTTP/1.1\r\nHost: example.com\r\n\r\n"; + assert_eq!(req_str.len(), 37); + let resp_str = "HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n"; + assert_eq!(resp_str.len(), 38); + + let parse_req = |max_header_size: Option| { + let mut bytes = BytesMut::from(req_str); + Server::parse( + &mut bytes, + ParseContext { + cached_headers: &mut None, + req_method: &mut None, + h1_parser_config: Default::default(), + h1_max_headers: None, + h1_max_header_size: max_header_size, + preserve_header_case: false, + #[cfg(feature = "ffi")] + preserve_header_order: false, + h09_responses: false, + #[cfg(feature = "client")] + on_informational: &mut None, + }, + ) + }; + + let parse_resp = |max_header_size: Option| { + let mut bytes = BytesMut::from(resp_str); + Client::parse( + &mut bytes, + ParseContext { + cached_headers: &mut None, + req_method: &mut None, + h1_parser_config: Default::default(), + h1_max_headers: None, + h1_max_header_size: max_header_size, + preserve_header_case: false, + #[cfg(feature = "ffi")] + preserve_header_order: false, + h09_responses: false, + #[cfg(feature = "client")] + on_informational: &mut None, + }, + ) + }; + + // Server checks + parse_req(None).unwrap().unwrap(); + parse_req(Some(37)).unwrap().unwrap(); + parse_req(Some(50)).unwrap().unwrap(); + assert!(matches!(parse_req(Some(36)), Err(Parse::TooLarge))); + assert!(matches!(parse_req(Some(10)), Err(Parse::TooLarge))); + + // Client checks + parse_resp(None).unwrap().unwrap(); + parse_resp(Some(38)).unwrap().unwrap(); + parse_resp(Some(50)).unwrap().unwrap(); + assert!(matches!(parse_resp(Some(37)), Err(Parse::TooLarge))); + assert!(matches!(parse_resp(Some(10)), Err(Parse::TooLarge))); + } + #[test] fn test_is_complete_fast() { let s = b"GET / HTTP/1.1\r\na: b\r\n\r\n"; @@ -3014,6 +3106,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -3097,6 +3190,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, @@ -3142,6 +3236,7 @@ mod tests { req_method: &mut None, h1_parser_config: Default::default(), h1_max_headers: None, + h1_max_header_size: None, preserve_header_case: false, #[cfg(feature = "ffi")] preserve_header_order: false, diff --git a/src/server/conn/http1.rs b/src/server/conn/http1.rs index 7bd3d72e2f..f89d0879ec 100644 --- a/src/server/conn/http1.rs +++ b/src/server/conn/http1.rs @@ -77,6 +77,7 @@ pub struct Builder { h1_title_case_headers: bool, h1_preserve_header_case: bool, h1_max_headers: Option, + h1_max_header_size: Option, h1_header_read_timeout: Dur, h1_writev: Option, max_buf_size: Option, @@ -247,6 +248,7 @@ impl Builder { h1_title_case_headers: false, h1_preserve_header_case: false, h1_max_headers: None, + h1_max_header_size: None, h1_header_read_timeout: Dur::Default(Some(Duration::from_secs(30))), h1_writev: None, max_buf_size: None, @@ -340,6 +342,17 @@ impl Builder { self } + /// Set the maximum size of request headers (including the start line) in bytes. + /// + /// If the client sends headers exceeding this limit, the server responds to the + /// client with "431 Request Header Fields Too Large" and closes the connection. + /// + /// Default is `None` (unbounded, or bounded by `max_buf_size`). + pub fn max_header_size(&mut self, val: usize) -> &mut Self { + self.h1_max_header_size = Some(val); + self + } + /// Set a timeout for reading client request headers. If a client does not /// transmit the entire header within this time, the connection is closed. /// @@ -475,6 +488,9 @@ impl Builder { if let Some(max_headers) = self.h1_max_headers { conn.set_http1_max_headers(max_headers); } + if let Some(max_header_size) = self.h1_max_header_size { + conn.set_http1_max_header_size(max_header_size); + } if let Some(dur) = self .timer .check(self.h1_header_read_timeout, "header_read_timeout") diff --git a/tests/client.rs b/tests/client.rs index b512260cc5..13c2a25889 100644 --- a/tests/client.rs +++ b/tests/client.rs @@ -2247,6 +2247,71 @@ mod conn { let _res = client.send_request(req).await.expect("send_request"); } + #[tokio::test] + async fn client_max_header_size_exceeded() { + let (server, addr) = setup_std_test_server(); + + thread::spawn(move || { + let mut sock = server.accept().unwrap().0; + let mut buf = [0; 1024]; + sock.read(&mut buf).unwrap(); + sock.write_all(b"HTTP/1.1 200 OK\r\nX-Long: ").unwrap(); + sock.write_all(&[b'a'; 1000]).unwrap(); + sock.write_all(b"\r\n\r\n").unwrap(); + }); + + let tcp = tcp_connect(&addr).await.unwrap(); + + let (mut client, conn) = conn::http1::Builder::new() + .max_header_size(512) + .handshake(tcp) + .await + .unwrap(); + + tokio::spawn(async move { + let _ = conn.await; + }); + + let req = Request::builder() + .uri("/a") + .body(Empty::::new()) + .unwrap(); + let err = client.send_request(req).await.unwrap_err(); + assert!(err.is_parse_too_large()); + } + + #[tokio::test] + async fn client_max_header_size_accepted() { + let (server, addr) = setup_std_test_server(); + + thread::spawn(move || { + let mut sock = server.accept().unwrap().0; + let mut buf = [0; 1024]; + sock.read(&mut buf).unwrap(); + sock.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n") + .unwrap(); + }); + + let tcp = tcp_connect(&addr).await.unwrap(); + + let (mut client, conn) = conn::http1::Builder::new() + .max_header_size(512) + .handshake(tcp) + .await + .unwrap(); + + tokio::spawn(async move { + let _ = conn.await; + }); + + let req = Request::builder() + .uri("/a") + .body(Empty::::new()) + .unwrap(); + let res = client.send_request(req).await.expect("send_request"); + assert_eq!(res.status(), StatusCode::OK); + } + #[tokio::test] async fn client_on_informational_ext() { use std::sync::atomic::{AtomicUsize, Ordering}; diff --git a/tests/server.rs b/tests/server.rs index 098ed7a6ee..094558cb61 100644 --- a/tests/server.rs +++ b/tests/server.rs @@ -2784,6 +2784,93 @@ async fn max_buf_size_split_header_boundary() { .expect_err("should TooLarge error"); } +#[cfg(feature = "http1")] +#[tokio::test] +async fn max_header_size_exceeded() { + let (listener, addr) = setup_tcp_listener(); + + const MAX_HEADER: usize = 512; + + thread::spawn(move || { + let mut tcp = connect(&addr); + tcp.write_all(b"GET / HTTP/1.1\r\nHost: x\r\nX-Long: ") + .expect("write 1"); + tcp.write_all(&[b'a'; 1000]).expect("write 2"); + tcp.write_all(b"\r\n\r\n").expect("write 3"); + let mut buf = [0; 256]; + tcp.read(&mut buf).expect("read 1"); + + let expected = "HTTP/1.1 431 "; + assert_eq!(s(&buf[..expected.len()]), expected); + }); + + let (socket, _) = listener.accept().await.unwrap(); + let socket = TokioIo::new(socket); + http1::Builder::new() + .max_header_size(MAX_HEADER) + .serve_connection(socket, HelloWorld) + .await + .expect_err("should TooLarge error"); +} + +#[cfg(feature = "http1")] +#[tokio::test] +async fn max_header_size_accepted() { + let (listener, addr) = setup_tcp_listener(); + + const MAX_HEADER: usize = 512; + + thread::spawn(move || { + let mut tcp = connect(&addr); + tcp.write_all(b"GET / HTTP/1.1\r\nHost: x\r\nConnection: close\r\n\r\n") + .expect("write"); + let mut buf = String::new(); + tcp.read_to_string(&mut buf).expect("read"); + + let expected = "HTTP/1.1 200 "; + assert_eq!(&buf[..expected.len()], expected); + }); + + let (socket, _) = listener.accept().await.unwrap(); + let socket = TokioIo::new(socket); + http1::Builder::new() + .max_header_size(MAX_HEADER) + .serve_connection(socket, HelloWorld) + .await + .expect("should succeed"); +} + +#[cfg(feature = "http1")] +#[tokio::test] +async fn max_header_size_with_large_max_buf_size() { + let (listener, addr) = setup_tcp_listener(); + + const MAX_HEADER: usize = 512; + const MAX_BUF: usize = 64 * 1024; + + thread::spawn(move || { + let mut tcp = connect(&addr); + tcp.write_all(b"GET / HTTP/1.1\r\nHost: x\r\nX-Long: ") + .expect("write 1"); + tcp.write_all(&[b'a'; 1000]).expect("write 2"); + tcp.write_all(b"\r\n\r\n").expect("write 3"); + let mut buf = [0; 256]; + tcp.read(&mut buf).expect("read 1"); + + let expected = "HTTP/1.1 431 "; + assert_eq!(s(&buf[..expected.len()]), expected); + }); + + let (socket, _) = listener.accept().await.unwrap(); + let socket = TokioIo::new(socket); + http1::Builder::new() + .max_buf_size(MAX_BUF) + .max_header_size(MAX_HEADER) + .serve_connection(socket, HelloWorld) + .await + .expect_err("should TooLarge error even with large max_buf_size"); +} + #[cfg(feature = "http1")] #[tokio::test] async fn graceful_shutdown_before_first_request_no_block() {