This repository has no description
0

Configure Feed

Select the types of activity you want to include in your feed.

core / knot2 / crates / knot-edge / src / quic.rs
18 kB 501 lines
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}