juicehost/
quic.rs

1//! QUIC/HTTP/3 server for juicehost.
2//!
3//! Accepts QUIC connections, wraps them in HTTP/3, and proxies requests
4//! through the same axum router that handles TCP traffic.
5
6use std::sync::Arc;
7
8use axum::body::Body;
9use axum::http::{Request, Response};
10use bytes::{Buf, Bytes};
11use h3_quinn::{quinn::Endpoint, Connection as H3Connection};
12use http_body_util::BodyExt;
13use quinn_proto::crypto::rustls::QuicServerConfig;
14use rcgen::generate_simple_self_signed;
15use rustls_pki_types::{CertificateDer, PrivateKeyDer};
16use tokio::sync::Notify;
17use tower::Service;
18
19use crate::server::build_router;
20use crate::state::AppState;
21
22fn generate_self_signed_cert() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
23    let certified_key = generate_simple_self_signed(vec!["juicebox.local".into()]).unwrap();
24    let cert_der = certified_key.cert.der().clone();
25    let key_der = PrivateKeyDer::try_from(certified_key.key_pair.serialize_der()).unwrap();
26    (cert_der, key_der)
27}
28
29/// Start the QUIC/HTTP/3 server on a UDP socket.
30///
31/// Wraps connections in HTTP/3 via h3 and proxies requests through the same
32/// axum router as the TCP listener. Shuts down when told to via the Notify.
33pub async fn start_quic_server(
34    state: Arc<AppState>,
35    addr: std::net::SocketAddr,
36    shutdown: Arc<Notify>,
37) {
38    let (cert_der, key_der) = generate_self_signed_cert();
39
40    let tls_config = {
41        let _ = rustls::crypto::aws_lc_rs::default_provider()
42            .install_default();
43        let mut tls = rustls::ServerConfig::builder()
44            .with_no_client_auth()
45            .with_single_cert(vec![cert_der], key_der)
46            .unwrap();
47        tls.alpn_protocols = vec![b"h3".to_vec()];
48        tls
49    };
50
51    let quic_server_config =
52        QuicServerConfig::try_from(Arc::new(tls_config)).expect("QuicServerConfig creation failed");
53    let server_config =
54        h3_quinn::quinn::ServerConfig::with_crypto(Arc::new(quic_server_config));
55
56    let endpoint =
57        Endpoint::server(server_config, addr).expect("Failed to bind QUIC endpoint");
58
59    tracing::info!("juicehost QUIC listening on udp://{}", addr);
60
61    let quic_shutdown = shutdown.clone();
62    tokio::select! {
63        biased;
64        _ = quic_shutdown.notified() => {
65            tracing::info!("juicehost QUIC shutting down...");
66            endpoint.wait_idle().await;
67        }
68        _ = run_quic_server(&endpoint, state) => {
69            endpoint.wait_idle().await;
70        }
71    }
72}
73
74async fn run_quic_server(endpoint: &Endpoint, state: Arc<AppState>) {
75    let router = build_router(Arc::clone(&state));
76
77    loop {
78        let conn = match endpoint.accept().await {
79            Some(conn) => conn,
80            None => break,
81        };
82
83        let router = router.clone();
84        tokio::spawn(async move {
85            let conn = match conn.await {
86                Ok(c) => c,
87                Err(e) => {
88                    tracing::warn!("QUIC connection handshake failed: {}", e);
89                    return;
90                }
91            };
92
93            if let Err(e) = handle_h3_connection(H3Connection::new(conn), router).await {
94                tracing::debug!("QUIC connection closed: {}", e);
95            }
96        });
97    }
98}
99
100async fn handle_h3_connection(
101    conn: H3Connection,
102    router: axum::Router,
103) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
104    let mut h3_conn = h3::server::builder().build(conn).await?;
105
106    loop {
107        let resolver = match h3_conn.accept().await {
108            Ok(Some(r)) => r,
109            Ok(None) => break,
110            Err(e) => {
111                tracing::debug!("h3 accept error: {}", e);
112                break;
113            }
114        };
115
116        let mut router = router.clone();
117        tokio::spawn(async move {
118            let (req, mut stream) = match resolver.resolve_request().await {
119                Ok(r) => r,
120                Err(e) => {
121                    tracing::debug!("h3 resolve_request error: {}", e);
122                    return;
123                }
124            };
125
126            if let Err(e) = proxy_axum(&mut router, req, &mut stream).await {
127                tracing::debug!("h3->axum proxy error: {}", e);
128            }
129        });
130    }
131
132    Ok(())
133}
134
135async fn proxy_axum(
136    router: &mut axum::Router,
137    req: Request<()>,
138    stream: &mut h3::server::RequestStream<h3_quinn::BidiStream<Bytes>, Bytes>,
139) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
140    let uri = req.uri().clone();
141    let method = req.method().clone();
142    let headers = req.headers().clone();
143
144    let mut body_bytes = Vec::new();
145    while let Some(chunk) = stream.recv_data().await? {
146        let chunk = chunk.chunk();
147        body_bytes.extend_from_slice(chunk);
148    }
149
150    let mut axum_req = Request::builder()
151        .method(method)
152        .uri(uri)
153        .body(Body::from(body_bytes))
154        .unwrap();
155    *axum_req.headers_mut() = headers;
156
157    let response = Service::call(router, axum_req).await?;
158
159    let status = response.status();
160    let resp_headers = response.headers().clone();
161    let resp_body = response.into_body();
162    let body_bytes: Vec<u8> = resp_body.collect().await?.to_bytes().to_vec();
163
164    let mut builder = Response::builder().status(status);
165    for (k, v) in resp_headers.iter() {
166        builder = builder.header(k, v);
167    }
168    stream
169        .send_response(builder.body(()).unwrap())
170        .await?;
171    stream.send_data(Bytes::from(body_bytes)).await?;
172    stream.finish().await?;
173
174    Ok(())
175}