This repository has no description
1use std::net::SocketAddr;
2use std::sync::Arc;
3use std::time::Duration;
4
5use axum::Router;
6use axum::body::Body;
7use axum::extract::ConnectInfo;
8use bytes::{Buf, Bytes};
9use futures::{FutureExt, StreamExt};
10use http::{Request, Response, Version};
11use quinn::{Endpoint, Incoming};
12use tokio::sync::Semaphore;
13use tokio_util::sync::CancellationToken;
14use tokio_util::task::TaskTracker;
15use tower::ServiceExt;
16
17use rustls::server::ResolvesServerCert;
18
19use crate::limits::ListenLimits;
20use crate::tls::{self, TlsError};
21use crate::zerortt::{EarlyData, EarlyDataPolicy};
22
23fn early_data(confirmed: bool) -> EarlyData {
24 match confirmed {
25 true => EarlyData::No,
26 false => EarlyData::Yes,
27 }
28}
29
30const LISTENER_DRAIN_GRACE: Duration = Duration::from_secs(30);
31const ENDPOINT_DRAIN_GRACE: Duration = Duration::from_secs(5);
32const CONNECTION_DRAIN_GRACE: Duration = Duration::from_secs(10);
33
34pub fn build_endpoint(
35 addr: SocketAddr,
36 resolver: Arc<dyn ResolvesServerCert>,
37 limits: ListenLimits,
38 early_data: EarlyDataPolicy,
39) -> Result<Endpoint, EndpointError> {
40 let server_config = tls::build_quic_server_config(resolver, limits, early_data)?;
41 let endpoint = Endpoint::server(server_config, addr)?;
42 Ok(endpoint)
43}
44
45#[derive(Debug, thiserror::Error)]
46pub enum EndpointError {
47 #[error(transparent)]
48 Tls(#[from] TlsError),
49 #[error("binding quic socket: {0}")]
50 Bind(#[from] std::io::Error),
51}
52
53pub async fn serve_http3(
54 endpoint: Endpoint,
55 app: Router,
56 limits: ListenLimits,
57 shutdown: CancellationToken,
58) {
59 let tracker = TaskTracker::new();
60 let connections = Arc::new(Semaphore::new(limits.max_connections()));
61 loop {
62 let incoming = tokio::select! {
63 () = shutdown.cancelled() => break,
64 incoming = endpoint.accept() => incoming,
65 };
66 let Some(incoming) = incoming else { break };
67 let Ok(permit) = Arc::clone(&connections).try_acquire_owned() else {
68 incoming.refuse();
69 continue;
70 };
71 let app = app.clone();
72 let conn_shutdown = shutdown.clone();
73 let conn_tracker = tracker.clone();
74 tracker.spawn(async move {
75 let _permit = permit;
76 if let Err(error) = serve_connection(incoming, app, conn_shutdown, conn_tracker).await {
77 tracing::debug!("h3 connection ended: {error}");
78 }
79 });
80 }
81 tracker.close();
82 let _ = tokio::time::timeout(LISTENER_DRAIN_GRACE, tracker.wait()).await;
83 endpoint.close(0u32.into(), b"shutdown");
84 let _ = tokio::time::timeout(ENDPOINT_DRAIN_GRACE, endpoint.wait_idle()).await;
85}
86
87async fn serve_connection(
88 incoming: Incoming,
89 app: Router,
90 shutdown: CancellationToken,
91 tracker: TaskTracker,
92) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
93 let (conn, confirmation, mut confirmed) = match incoming.accept()?.into_0rtt() {
94 Ok((conn, accepted)) => (conn, accepted.map(|_| ()).left_future(), false),
95 Err(connecting) => (
96 connecting.await?,
97 std::future::pending::<()>().right_future(),
98 true,
99 ),
100 };
101 tokio::pin!(confirmation);
102 let remote = conn.remote_address();
103 let quic = conn.clone();
104 tracing::trace!("h3 connection from {remote} accepted");
105 let mut h3_conn =
106 h3::server::Connection::<_, Bytes>::new(h3_quinn::Connection::new(conn)).await?;
107 tracing::trace!("h3 connection from {remote} established");
108
109 let drain_deadline = tokio::time::sleep(CONNECTION_DRAIN_GRACE);
110 tokio::pin!(drain_deadline);
111 let mut draining = false;
112 loop {
113 tokio::select! {
114 biased;
115 () = shutdown.cancelled(), if !draining => {
116 draining = true;
117 drain_deadline
118 .as_mut()
119 .reset(tokio::time::Instant::now() + CONNECTION_DRAIN_GRACE);
120 let _ = h3_conn.shutdown(0).await;
121 }
122 () = &mut drain_deadline, if draining => {
123 tracing::debug!("h3 connection from {remote} drain timed out, closing");
124 quic.close(0u32.into(), b"drain timeout");
125 break;
126 }
127 resolved = h3_conn.accept() => match resolved {
128 Ok(Some(resolver)) => {
129 let app = app.clone();
130 let early = early_data(confirmed);
131 tracker.spawn(async move {
132 if let Err(error) = serve_request(resolver, app, remote, early).await {
133 tracing::debug!("h3 request from {remote} failed: {error}");
134 }
135 });
136 }
137 Ok(None) => {
138 tracing::debug!("h3 connection from {remote} closed by the client");
139 break;
140 }
141 Err(error) => {
142 tracing::debug!("h3 accept from {remote} error: {error}");
143 break;
144 }
145 },
146 () = &mut confirmation, if !confirmed => {
147 confirmed = true;
148 }
149 }
150 }
151 Ok(())
152}
153
154async fn serve_request(
155 resolver: h3::server::RequestResolver<h3_quinn::Connection, Bytes>,
156 app: Router,
157 remote: SocketAddr,
158 early: EarlyData,
159) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
160 let (request, stream) = resolver.resolve_request().await?;
161 let (mut send, recv) = stream.split();
162
163 let (mut parts, ()) = request.into_parts();
164 parts.version = Version::HTTP_3;
165 parts.extensions.insert(ConnectInfo(remote));
166 parts.extensions.insert(early);
167 let request = Request::from_parts(parts, request_body(recv));
168
169 let response = match app.oneshot(request).await {
170 Ok(response) => response,
171 Err(infallible) => match infallible {},
172 };
173
174 let (parts, body) = response.into_parts();
175 send.send_response(Response::from_parts(parts, ())).await?;
176
177 let mut data = body.into_data_stream();
178 while let Some(chunk) = data.next().await {
179 match chunk {
180 Ok(bytes) if bytes.has_remaining() => send.send_data(bytes).await?,
181 Ok(_) => {}
182 Err(error) => {
183 tracing::debug!("h3 response body to {remote} errored: {error}");
184 send.stop_stream(h3::error::Code::H3_INTERNAL_ERROR);
185 return Ok(());
186 }
187 }
188 }
189 send.finish().await?;
190 Ok(())
191}
192
193struct RecvGuard {
194 stream: h3::server::RequestStream<h3_quinn::RecvStream, Bytes>,
195 ended: bool,
196}
197
198impl Drop for RecvGuard {
199 fn drop(&mut self) {
200 if !self.ended {
201 self.stream.stop_sending(h3::error::Code::H3_NO_ERROR);
202 }
203 }
204}
205
206fn request_body(recv: h3::server::RequestStream<h3_quinn::RecvStream, Bytes>) -> Body {
207 let guard = RecvGuard {
208 stream: recv,
209 ended: false,
210 };
211 let stream = futures::stream::unfold(Some(guard), |state| async move {
212 let mut guard = state?;
213 match guard.stream.recv_data().await {
214 Ok(Some(mut buf)) => {
215 let bytes = buf.copy_to_bytes(buf.remaining());
216 Some((Ok::<Bytes, std::io::Error>(bytes), Some(guard)))
217 }
218 Ok(None) => {
219 guard.ended = true;
220 None
221 }
222 Err(error) => {
223 guard.ended = true;
224 Some((Err(std::io::Error::other(error.to_string())), None))
225 }
226 }
227 });
228 Body::from_stream(stream)
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234
235 use axum::routing::get;
236 use rustls::crypto::aws_lc_rs;
237
238 use crate::tls;
239
240 #[test]
241 fn an_unconfirmed_handshake_is_early_data_and_a_confirmed_one_is_not() {
242 assert_eq!(
243 early_data(false),
244 EarlyData::Yes,
245 "data before handshake confirmation is early data, fail closed"
246 );
247 assert_eq!(early_data(true), EarlyData::No);
248 }
249
250 fn client_endpoint() -> Endpoint {
251 client_endpoint_with_provider(aws_lc_rs::default_provider())
252 }
253
254 fn client_endpoint_with_provider(provider: rustls::crypto::CryptoProvider) -> Endpoint {
255 let mut crypto = rustls::ClientConfig::builder_with_provider(Arc::new(provider))
256 .with_protocol_versions(&[&rustls::version::TLS13])
257 .unwrap()
258 .dangerous()
259 .with_custom_certificate_verifier(Arc::new(tls::test_support::AcceptAnyServerCert))
260 .with_no_client_auth();
261 crypto.alpn_protocols = vec![b"h3".to_vec()];
262 let quic = quinn::crypto::rustls::QuicClientConfig::try_from(crypto).unwrap();
263 let mut endpoint = Endpoint::client("[::1]:0".parse().unwrap()).unwrap();
264 endpoint.set_default_client_config(quinn::ClientConfig::new(Arc::new(quic)));
265 endpoint
266 }
267
268 fn spawn_h3_server(
269 app: Router,
270 limits: ListenLimits,
271 early_data: EarlyDataPolicy,
272 ) -> (SocketAddr, CancellationToken) {
273 let endpoint = build_endpoint(
274 "[::1]:0".parse().unwrap(),
275 tls::test_support::resolver(),
276 limits,
277 early_data,
278 )
279 .unwrap();
280 let addr = endpoint.local_addr().unwrap();
281 let shutdown = CancellationToken::new();
282 tokio::spawn(serve_http3(
283 endpoint,
284 crate::altsvc::with_host_from_authority(app),
285 limits,
286 shutdown.clone(),
287 ));
288 (addr, shutdown)
289 }
290
291 async fn h3_client_connect(
292 client: &Endpoint,
293 addr: SocketAddr,
294 ) -> h3::client::SendRequest<h3_quinn::OpenStreams, Bytes> {
295 let conn = client.connect(addr, "localhost").unwrap().await.unwrap();
296 let (mut driver, send_request) = h3::client::new(h3_quinn::Connection::new(conn))
297 .await
298 .unwrap();
299 tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
300 send_request
301 }
302
303 async fn h3_request(
304 send_request: &mut h3::client::SendRequest<h3_quinn::OpenStreams, Bytes>,
305 path: &str,
306 ) -> (http::StatusCode, Vec<u8>) {
307 let request = http::Request::get(format!("https://localhost{path}"))
308 .body(())
309 .unwrap();
310 let mut stream = send_request.send_request(request).await.unwrap();
311 stream.finish().await.unwrap();
312 let status = stream.recv_response().await.unwrap().status();
313 let chunks: Vec<Bytes> = futures::stream::unfold(stream, |mut stream| async move {
314 stream
315 .recv_data()
316 .await
317 .unwrap()
318 .map(|mut buf| (buf.copy_to_bytes(buf.remaining()), stream))
319 })
320 .collect()
321 .await;
322 let body = chunks
323 .iter()
324 .flat_map(|chunk| chunk.iter().copied())
325 .collect();
326 (status, body)
327 }
328
329 #[tokio::test]
330 async fn an_h3_get_roundtrips_through_the_router() {
331 let app = Router::new().route(
332 "/",
333 get(|ConnectInfo(peer): ConnectInfo<SocketAddr>| async move { peer.to_string() }),
334 );
335 let (addr, shutdown) =
336 spawn_h3_server(app, tls::test_support::limits(), EarlyDataPolicy::Disabled);
337
338 let client = client_endpoint();
339 let client_port = client.local_addr().unwrap().port();
340 let mut send_request = h3_client_connect(&client, addr).await;
341 let (status, body) = h3_request(&mut send_request, "/").await;
342 assert_eq!(status, 200);
343 let reported: SocketAddr = String::from_utf8(body).unwrap().parse().unwrap();
344 assert_eq!(
345 reported.port(),
346 client_port,
347 "the h3 handler must see the QUIC remote address via ConnectInfo"
348 );
349 shutdown.cancel();
350 }
351
352 #[tokio::test]
353 async fn an_h3_request_is_tagged_with_the_h3_protocol() {
354 use axum::middleware::from_fn;
355
356 let app = Router::new()
357 .route(
358 "/proto",
359 get(|req: axum::extract::Request| async move {
360 req.extensions()
361 .get::<crate::protocol::NegotiatedProtocol>()
362 .map(|protocol| protocol.as_str())
363 .unwrap_or("missing")
364 .to_string()
365 }),
366 )
367 .layer(from_fn(crate::protocol::tag));
368 let (addr, shutdown) =
369 spawn_h3_server(app, tls::test_support::limits(), EarlyDataPolicy::Disabled);
370
371 let client = client_endpoint();
372 let mut send_request = h3_client_connect(&client, addr).await;
373 let (status, body) = h3_request(&mut send_request, "/proto").await;
374 assert_eq!(status, 200);
375 assert_eq!(
376 String::from_utf8(body).unwrap(),
377 "h3",
378 "a request served over QUIC must be tagged as the h3 negotiated protocol"
379 );
380 shutdown.cancel();
381 }
382
383 #[tokio::test]
384 async fn an_early_data_enabled_endpoint_still_serves_through_the_zero_rtt_path() {
385 let app = Router::new().route("/", get(|| async { "ok" }));
386 let (addr, shutdown) =
387 spawn_h3_server(app, tls::test_support::limits(), EarlyDataPolicy::Enabled);
388
389 let client = client_endpoint();
390 let mut send_request = h3_client_connect(&client, addr).await;
391 let (status, _) = h3_request(&mut send_request, "/").await;
392 assert_eq!(
393 status, 200,
394 "a server with early data enabled must serve requests through the 0-RTT acceptance path"
395 );
396 shutdown.cancel();
397 }
398
399 #[tokio::test]
400 async fn a_classical_only_h3_client_completes_over_x25519() {
401 let app = Router::new().route("/", get(|| async { "ok" }));
402 let (addr, shutdown) =
403 spawn_h3_server(app, tls::test_support::limits(), EarlyDataPolicy::Disabled);
404
405 let client = client_endpoint_with_provider(tls::test_support::classical_only_provider());
406 let mut send_request = h3_client_connect(&client, addr).await;
407 let (status, _) = h3_request(&mut send_request, "/").await;
408 assert_eq!(
409 status, 200,
410 "a QUIC client without ML-KEM must still complete the h3 handshake over classical X25519"
411 );
412 shutdown.cancel();
413 }
414
415 #[tokio::test]
416 async fn an_in_flight_h3_request_finishes_during_drain() {
417 let app = Router::new().route(
418 "/slow",
419 get(|| async {
420 tokio::time::sleep(Duration::from_millis(300)).await;
421 "drained-clean"
422 }),
423 );
424 let (addr, shutdown) =
425 spawn_h3_server(app, tls::test_support::limits(), EarlyDataPolicy::Disabled);
426
427 let client = client_endpoint();
428 let mut send_request = h3_client_connect(&client, addr).await;
429 let ((status, body), ()) = tokio::join!(h3_request(&mut send_request, "/slow"), async {
430 tokio::time::sleep(Duration::from_millis(100)).await;
431 shutdown.cancel();
432 });
433 assert_eq!(
434 status, 200,
435 "an in-flight h3 request must complete through the graceful drain"
436 );
437 assert_eq!(body, b"drained-clean");
438 shutdown.cancel();
439 }
440
441 #[tokio::test]
442 async fn an_h3_connection_survives_repeated_client_path_migration() {
443 let app = Router::new().route("/echo", get(|| async { "migrated-clean" }));
444 let (addr, shutdown) =
445 spawn_h3_server(app, tls::test_support::limits(), EarlyDataPolicy::Disabled);
446
447 let client = client_endpoint();
448 let send_request = h3_client_connect(&client, addr).await;
449 futures::stream::iter(0u8..3)
450 .fold(
451 (client, send_request),
452 |(client, mut send_request), round| async move {
453 client
454 .rebind(std::net::UdpSocket::bind("[::1]:0").unwrap())
455 .unwrap();
456 let (status, body) = h3_request(&mut send_request, "/echo").await;
457 assert_eq!(
458 status, 200,
459 "round {round}: a request after path migration must still be served"
460 );
461 assert_eq!(
462 body, b"migrated-clean",
463 "round {round}: the migrated connection must deliver the response uncorrupted"
464 );
465 (client, send_request)
466 },
467 )
468 .await;
469
470 shutdown.cancel();
471 }
472
473 #[tokio::test]
474 async fn a_slow_h3_response_survives_the_idle_timeout() {
475 use std::num::{NonZeroU32, NonZeroU64};
476
477 let limits = ListenLimits::new(
478 crate::limits::HeaderTimeout::from_millis(NonZeroU64::new(1_000).unwrap()),
479 crate::limits::IdleTimeout::from_millis(NonZeroU64::new(2_000).unwrap()),
480 NonZeroU32::new(64).unwrap(),
481 );
482 let app = Router::new().route(
483 "/",
484 get(|| async {
485 tokio::time::sleep(Duration::from_secs(3)).await;
486 "ok"
487 }),
488 );
489 let (addr, shutdown) = spawn_h3_server(app, limits, EarlyDataPolicy::Disabled);
490
491 let client = client_endpoint();
492 let mut send_request = h3_client_connect(&client, addr).await;
493 let (status, body) = h3_request(&mut send_request, "/").await;
494 assert_eq!(
495 status, 200,
496 "keep-alive must hold the connection through a response slower than the idle timeout"
497 );
498 assert_eq!(body, b"ok");
499 shutdown.cancel();
500 }
501}