diff --git a/src/proto/h2/mod.rs b/src/proto/h2/mod.rs index 393d4179ab..506b8bc407 100644 --- a/src/proto/h2/mod.rs +++ b/src/proto/h2/mod.rs @@ -59,8 +59,9 @@ fn strip_connection_headers(headers: &mut HeaderMap, kind: MessageKind) { #[cfg(feature = "client")] if matches!(kind, MessageKind::Request) { if headers - .get(http::header::TE) - .map_or(false, |te_header| te_header != "trailers") + .get_all(http::header::TE) + .iter() + .any(|te_header| te_header != "trailers") { warn!("TE headers not set to \"trailers\" are illegal in HTTP/2 requests"); headers.remove(http::header::TE); diff --git a/tests/integration.rs b/tests/integration.rs index 2deee443f8..8e051254de 100644 --- a/tests/integration.rs +++ b/tests/integration.rs @@ -232,6 +232,54 @@ t! { ; } +#[tokio::test] +async fn h2_strips_te_header_with_repeated_values() { + use http_body_util::Empty; + use hyper::body::Bytes; + use tokio::net::{TcpListener, TcpStream}; + + let listener = TcpListener::bind(SocketAddr::from(([127, 0, 0, 1], 0))) + .await + .unwrap(); + let addr = listener.local_addr().unwrap(); + + tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let service = hyper::service::service_fn(|req: hyper::Request| { + let te_count = req.headers().get_all("te").iter().count(); + async move { + let mut res = hyper::Response::new(Empty::::new()); + res.headers_mut().insert("x-te-count", te_count.into()); + Ok::<_, hyper::Error>(res) + } + }); + hyper::server::conn::http2::Builder::new(TokioExecutor) + .serve_connection(TokioIo::new(stream), service) + .await + .unwrap(); + }); + + let stream = TcpStream::connect(addr).await.unwrap(); + let (mut client, conn) = hyper::client::conn::http2::Builder::new(TokioExecutor) + .handshake(TokioIo::new(stream)) + .await + .unwrap(); + tokio::spawn(conn); + + // Only the first TE value used to be checked, so the second one reached + // the peer and the request was rejected as malformed. + let req = hyper::Request::builder() + .uri(format!("http://{addr}/")) + .header("te", "trailers") + .header("te", "gzip") + .body(Empty::::new()) + .unwrap(); + let res = client.send_request(req).await.unwrap(); + + assert_eq!(res.status(), 200); + assert_eq!(res.headers()["x-te-count"], "0"); +} + t! { get_body_chunked, client: