1use bytes::{Buf, Bytes};
2use hyper::{Request, Response, StatusCode};
3use std::sync::Arc;
4
5use crate::server::{REQUEST_TIMEOUT, Server};
6
7impl<S> Server<S>
8where
9 S: Clone + Send + Sync + 'static,
10{
11 pub async fn serve_h3(self, mut quic_server: s2n_quic::Server) -> Result<(), std::io::Error> {
19 crate::server::enforce_fips_compliance()?;
20 let state = Arc::new(self);
21 let connection_semaphore = Arc::new(tokio::sync::Semaphore::new(state.max_connections));
22
23 while let Some(conn) = quic_server.accept().await {
24 let Ok(permit) = connection_semaphore.clone().acquire_owned().await else {
25 break;
26 };
27 let state = state.clone();
28 tokio::spawn(async move {
29 state.handle_h3_connection(conn).await;
30 drop(permit);
31 });
32 }
33 Ok(())
34 }
35
36 async fn handle_h3_connection(self: Arc<Self>, conn: s2n_quic::Connection) {
37 let Ok(peer) = conn.remote_addr() else {
38 return;
39 };
40
41 let h3_conn = s2n_quic_h3::Connection::new(conn);
42 let Ok(mut h3_server) = h3::server::Connection::new(h3_conn).await else {
43 return;
44 };
45
46 let stream_semaphore = Arc::new(tokio::sync::Semaphore::new(256));
48
49 loop {
50 let Ok(stream_permit) = stream_semaphore.clone().acquire_owned().await else {
51 break;
52 };
53 match h3_server.accept().await {
54 Ok(Some(resolver)) => {
55 let state = self.clone();
56 tokio::spawn(async move {
57 state.handle_h3_request(resolver, peer).await;
58 drop(stream_permit);
59 });
60 }
61 Ok(None) => {
62 drop(stream_permit);
63 break;
64 }
65 Err(e) => {
66 drop(stream_permit);
67 let err_str = e.to_string();
68 if !err_str.contains("application error")
69 && !err_str.contains("ConnectionError")
70 {
71 tracing::debug!("[h3] Stream accept error: {}", e);
72 }
73 break;
74 }
75 }
76 }
77 }
78
79 async fn read_h3_body(
80 &self,
81 parts: &hyper::http::request::Parts,
82 stream: &mut h3::server::RequestStream<s2n_quic_h3::BidiStream<Bytes>, Bytes>,
83 ) -> Result<Bytes, StatusCode> {
84 let method = &parts.method;
85 if method == hyper::Method::GET || method == hyper::Method::HEAD {
86 return Ok(Bytes::new());
87 }
88
89 let content_length = parts
90 .headers
91 .get(hyper::header::CONTENT_LENGTH)
92 .and_then(|v| v.to_str().ok())
93 .and_then(|s| s.parse::<usize>().ok())
94 .unwrap_or(0);
95
96 let cap = if content_length > 0 && content_length <= self.max_body_size {
97 content_length
98 } else {
99 0
100 };
101
102 let initial_allocation = if cap > 0 && cap <= 256 * 1024 {
105 cap
106 } else {
107 std::cmp::min(cap, 64 * 1024)
108 };
109 let mut body_vec = Vec::with_capacity(initial_allocation);
110
111 let timeout_res = tokio::time::timeout(REQUEST_TIMEOUT, async {
112 loop {
113 match stream.recv_data().await {
114 Ok(Some(mut chunk)) => {
115 while chunk.has_remaining() {
116 let data = chunk.chunk();
117 if body_vec.len() + data.len() > self.max_body_size {
118 return Err(StatusCode::PAYLOAD_TOO_LARGE);
119 }
120 body_vec.extend_from_slice(data);
121 let len = data.len();
122 chunk.advance(len);
123 }
124 }
125 Ok(None) => break,
126 Err(_) => return Err(StatusCode::BAD_REQUEST),
127 }
128 }
129 Ok(())
130 })
131 .await;
132
133 match timeout_res {
134 Ok(Err(status)) => Err(status),
135 Err(_) => Err(StatusCode::REQUEST_TIMEOUT),
136 Ok(Ok(())) => Ok(Bytes::from(body_vec)),
137 }
138 }
139
140 async fn handle_h3_request(
141 self: Arc<Self>,
142 resolver: h3::server::RequestResolver<s2n_quic_h3::Connection, Bytes>,
143 peer: std::net::SocketAddr,
144 ) {
145 let resolve_res = tokio::time::timeout(REQUEST_TIMEOUT, resolver.resolve_request()).await;
146
147 let Ok(Ok((req, mut stream))) = resolve_res else {
148 return;
149 };
150
151 let (parts, ()) = req.into_parts();
152
153 let body_bytes = match self.read_h3_body(&parts, &mut stream).await {
154 Ok(bytes) => bytes,
155 Err(status) => {
156 let _ = stream
157 .send_response(
158 Response::builder()
159 .status(status)
160 .body(())
161 .unwrap_or_else(|_| Response::new(())),
162 )
163 .await;
164 let _ = stream.finish().await;
165 return;
166 }
167 };
168
169 let mut rebuild_req =
176 Request::from_parts(parts, crate::http::response::Body::full(body_bytes));
177 #[cfg(feature = "original-uri")]
178 {
179 let orig_uri = rebuild_req.uri().clone();
180 rebuild_req
181 .extensions_mut()
182 .insert(crate::routing::extract::OriginalUri(orig_uri));
183 }
184 rebuild_req
185 .extensions_mut()
186 .insert(crate::routing::extract::ConnectInfo(peer));
187 rebuild_req
188 .extensions_mut()
189 .insert(crate::routing::extract::MaxBodySize(self.max_body_size));
190
191 let full_resp = self.router.handle_request(rebuild_req).await;
192
193 let (resp_parts, body) = full_resp.into_parts();
194 let resp = Response::from_parts(resp_parts, ());
195
196 if stream.send_response(resp).await.is_ok() {
197 use http_body_util::BodyExt;
198 let mut body = body;
199 while let Some(frame_res) = body.frame().await {
200 if let Ok(frame) = frame_res {
201 let send_res = if let Some(data) = frame.data_ref() {
202 stream.send_data(data.clone()).await
203 } else if let Some(trailers) = frame.trailers_ref() {
204 stream.send_trailers(trailers.clone()).await
205 } else {
206 Ok(())
207 };
208
209 if send_res.is_err() {
210 break;
211 }
212 } else {
213 break;
214 }
215 }
216 }
217 let _ = stream.finish().await;
218 }
219}
220
221#[cfg(all(test, feature = "cert-gen"))]
222mod tests {
223 #![allow(clippy::unwrap_used, clippy::expect_used)]
224
225 use crate::routing::{Router, get, post};
226 use crate::server::Server;
227 use bytes::{Buf, Bytes};
228 use rustls::pki_types::{CertificateDer, PrivateKeyDer};
229 use std::sync::Arc;
230
231 async fn hello() -> &'static str {
232 "hello from h3"
233 }
234
235 async fn echo(body: Bytes) -> Vec<u8> {
236 body.to_vec()
237 }
238
239 fn build_server_config(cert_pem: &str, key_pem: &str) -> rustls::ServerConfig {
244 let cert_chain: Vec<CertificateDer<'static>> =
245 rustls_pemfile::certs(&mut cert_pem.as_bytes())
246 .filter_map(Result::ok)
247 .collect();
248 let key_der: PrivateKeyDer<'static> = rustls_pemfile::private_key(&mut key_pem.as_bytes())
249 .expect("parse private key")
250 .expect("private key present in PEM");
251
252 let mut config = rustls::ServerConfig::builder()
253 .with_no_client_auth()
254 .with_single_cert(cert_chain, key_der)
255 .expect("build rustls ServerConfig");
256 config.alpn_protocols = vec![b"h3".to_vec()];
257 config
258 }
259
260 fn start_h3_server(
263 app: Router<()>,
264 max_body_size: Option<usize>,
265 ) -> (std::net::SocketAddr, String) {
266 let cert = crate::tls::generate_self_signed_cert(vec!["localhost".to_string()])
267 .expect("generate self-signed cert");
268 let config = build_server_config(&cert.cert_pem, &cert.key_pem);
269
270 let quic_tls = s2n_quic::provider::tls::rustls::Server::from(Arc::new(config));
271 let quic_server = s2n_quic::Server::builder()
272 .with_tls(quic_tls)
273 .expect("with_tls")
274 .with_io("127.0.0.1:0")
275 .expect("with_io")
276 .start()
277 .expect("start quic server");
278 let addr = quic_server.local_addr().expect("local addr");
279
280 let mut server = Server::new(app);
281 if let Some(limit) = max_body_size {
282 server = server.max_body_size(limit);
283 }
284 drop(tokio::spawn(async move {
285 let _ = server.serve_h3(quic_server).await;
286 }));
287
288 (addr, cert.cert_pem)
289 }
290
291 async fn h3_connect(
296 addr: std::net::SocketAddr,
297 cert_pem: &str,
298 ) -> (
299 h3::client::SendRequest<s2n_quic_h3::OpenStreams, Bytes>,
300 tokio::task::JoinHandle<()>,
301 ) {
302 let client_tls = s2n_quic::provider::tls::rustls::Client::builder()
303 .with_certificate(cert_pem)
304 .expect("with_certificate")
305 .with_application_protocols(std::iter::once("h3"))
306 .expect("with_application_protocols")
307 .build()
308 .expect("build client tls");
309 let client = s2n_quic::Client::builder()
310 .with_tls(client_tls)
311 .expect("with_tls")
312 .with_io("127.0.0.1:0")
313 .expect("with_io")
314 .start()
315 .expect("start quic client");
316
317 let quic_conn = client
318 .connect(s2n_quic::client::Connect::new(addr).with_server_name("localhost"))
319 .await
320 .expect("quic connect");
321
322 let h3_conn = s2n_quic_h3::Connection::new(quic_conn);
323 let (mut driver, send_request) = h3::client::new(h3_conn).await.expect("h3 client new");
324 let driver_task = tokio::spawn(async move {
325 let _ = driver.wait_idle().await;
326 });
327
328 (send_request, driver_task)
329 }
330
331 async fn recv_all<S>(stream: &mut h3::client::RequestStream<S, Bytes>) -> Vec<u8>
333 where
334 S: h3::quic::RecvStream,
335 {
336 let mut body = Vec::new();
337 while let Some(mut chunk) = stream.recv_data().await.expect("recv_data") {
338 while chunk.has_remaining() {
339 let n = chunk.remaining();
340 body.extend_from_slice(&chunk.copy_to_bytes(n));
341 }
342 }
343 body
344 }
345
346 #[tokio::test]
352 async fn h3_get_and_post_round_trip() {
353 let app = Router::new()
354 .route("/", get(hello))
355 .route("/echo", post(echo));
356 let (addr, cert_pem) = start_h3_server(app, None);
357
358 let (mut send_request, driver_task) = h3_connect(addr, &cert_pem).await;
359
360 let get_req = hyper::Request::builder()
362 .method("GET")
363 .uri("https://localhost/")
364 .body(())
365 .expect("build GET request");
366 let mut get_stream = send_request
367 .send_request(get_req)
368 .await
369 .expect("send GET request");
370 get_stream
371 .finish()
372 .await
373 .expect("finish GET request stream");
374 let get_response = get_stream.recv_response().await.expect("recv GET response");
375 assert_eq!(get_response.status(), hyper::StatusCode::OK);
376 let get_body = recv_all(&mut get_stream).await;
377 assert_eq!(get_body, b"hello from h3");
378
379 let payload = b"round trip me over quic".to_vec();
382 let post_req = hyper::Request::builder()
383 .method("POST")
384 .uri("https://localhost/echo")
385 .header(hyper::header::CONTENT_LENGTH, payload.len())
386 .body(())
387 .expect("build POST request");
388 let mut post_stream = send_request
389 .send_request(post_req)
390 .await
391 .expect("send POST request");
392 post_stream
393 .send_data(Bytes::from(payload.clone()))
394 .await
395 .expect("send POST body");
396 post_stream
397 .finish()
398 .await
399 .expect("finish POST request stream");
400 let post_response = post_stream
401 .recv_response()
402 .await
403 .expect("recv POST response");
404 assert_eq!(post_response.status(), hyper::StatusCode::OK);
405 let post_body = recv_all(&mut post_stream).await;
406 assert_eq!(post_body, payload);
407
408 drop(send_request);
409 driver_task.abort();
410 }
411
412 #[tokio::test]
415 async fn h3_post_over_max_body_size_is_rejected() {
416 let app = Router::new().route("/echo", post(echo));
417 let (addr, cert_pem) = start_h3_server(app, Some(8));
418
419 let (mut send_request, driver_task) = h3_connect(addr, &cert_pem).await;
420
421 let payload = vec![b'x'; 64];
422 let req = hyper::Request::builder()
423 .method("POST")
424 .uri("https://localhost/echo")
425 .header(hyper::header::CONTENT_LENGTH, payload.len())
426 .body(())
427 .expect("build POST request");
428 let mut stream = send_request
429 .send_request(req)
430 .await
431 .expect("send POST request");
432 stream
433 .send_data(Bytes::from(payload))
434 .await
435 .expect("send POST body");
436 stream.finish().await.expect("finish POST request stream");
437 let response = stream.recv_response().await.expect("recv response");
438 assert_eq!(response.status(), hyper::StatusCode::PAYLOAD_TOO_LARGE);
439
440 drop(send_request);
441 driver_task.abort();
442 }
443}