Skip to main content

tachyon_web/server/
h3.rs

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    /// Serve HTTP/3 over QUIC using the given s2n-quic server.
12    ///
13    /// # Errors
14    ///
15    /// Returns an error if FIPS compliance enforcement fails. The accept loop
16    /// itself never surfaces per-connection errors as an `Err`; it just stops
17    /// when `quic_server.accept()` returns `None`.
18    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        // Limit concurrent streams per connection for DoS protection.
47        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        // Strictly enforce zero-trust client allocation limits (DoS prevention).
103        // Pre-allocate up to 256 KiB directly to avoid reallocations for standard payloads.
104        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        // HTTP/3 stream frames aren't a standard `hyper::body::Body`, so (unlike the
170        // HTTP/1.1 and HTTP/2 paths) this body is fully buffered up front by
171        // `read_h3_body` rather than streamed lazily. It's still wrapped in the same
172        // unified `Body` type so it flows through the same extractor pipeline —
173        // `BodyStream`/`Request<Body>` handlers work identically, they just don't get
174        // the zero-buffering benefit over HTTP/3 in this version.
175        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    /// Builds a self-signed `rustls::ServerConfig` (ALPN "h3") from PEM cert/key strings.
240    /// Hand-rolled with `rustls_pemfile` (mirroring `Server::start_all_inner`/`RustlsConfig::
241    /// from_pem` in `server/mod.rs`) rather than reusing `TlsPolicy::server_config_from_pem`, so
242    /// this test module only needs `cert-gen` — not also `tor`/`i2p` — to compile.
243    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    /// Starts a real `Server::serve_h3` on loopback (OS-assigned port) serving `app`. Returns the
261    /// bound address and the self-signed cert's PEM (for the client to trust).
262    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    /// Connects a real HTTP/3 client (over loopback UDP) to `addr`, trusting `cert_pem`. The
292    /// returned `JoinHandle` drives the connection's control/QPACK streams in the background —
293    /// per `h3::client::Connection`'s own docs, this must stay alive and polled for the
294    /// connection to make progress while requests are in flight.
295    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    /// Reads all remaining `DATA` frames off a response stream into a `Vec<u8>`.
332    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    /// Full loopback HTTP/3 round trip: a real `s2n-quic`/`h3` client speaking QUIC to a real
347    /// `Server::serve_h3`. Exercises `handle_h3_connection`'s accept loop, `read_h3_body`'s
348    /// GET early-return and its `recv_data`/Content-Length-driven accumulation for POST, and
349    /// `handle_h3_request`'s full response path (`send_response`, the `frame()`/`send_data()`
350    /// loop, and `finish()`).
351    #[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        // GET / — covers the GET/HEAD early-return in `read_h3_body`.
361        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        // POST /echo with a Content-Length under `max_body_size` — covers the `recv_data`
380        // accumulation loop and the Content-Length pre-allocation branch in `read_h3_body`.
381        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    /// A POST body exceeding `max_body_size` — covers the `PAYLOAD_TOO_LARGE` branch in
413    /// `read_h3_body`'s `recv_data` loop.
414    #[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}