juiceback/
quic.rs

1//! QUIC/HTTP/3 server and client for juiceback.
2//!
3//! The server side accepts QUIC connections, wraps them in HTTP/3, and proxies
4//! requests through the same axum router used for TCP. The client side pushes
5//! file data to juicehost over QUIC instead of regular TCP.
6
7use std::sync::Arc;
8
9use axum::body::Body;
10use axum::http::{Request, Response, Uri};
11
12use futures::future;
13use rustls::client::danger::{
14    HandshakeSignatureValid, ServerCertVerifier,
15    ServerCertVerified,
16};
17use rustls::{DigitallySignedStruct, Error as RustlsError, SignatureScheme};
18use bytes::{Buf, Bytes};
19use h3_quinn::{quinn::Endpoint, Connection as H3Connection};
20use http_body_util::BodyExt;
21use quinn_proto::crypto::rustls::QuicServerConfig;
22use rcgen::generate_simple_self_signed;
23use rustls_pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime};
24use tokio::sync::Notify;
25use tower::Service;
26
27use crate::routes;
28use crate::state::AppState;
29
30// -- Server --
31
32/// Generate an ephemeral self-signed TLS certificate for QUIC.
33fn generate_self_signed_cert() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
34    let certified_key = generate_simple_self_signed(vec!["juicebox.local".into()]).unwrap();
35    let cert_der = certified_key.cert.der().clone();
36    let key_der = PrivateKeyDer::try_from(certified_key.key_pair.serialize_der()).unwrap();
37    (cert_der, key_der)
38}
39
40pub async fn start_quic_server(
41    state: Arc<AppState>,
42    addr: std::net::SocketAddr,
43    shutdown: Arc<Notify>,
44) {
45    let (cert_der, key_der) = generate_self_signed_cert();
46
47    let tls_config = {
48        let _ = rustls::crypto::aws_lc_rs::default_provider()
49            .install_default();
50        let mut tls = rustls::ServerConfig::builder()
51            .with_no_client_auth()
52            .with_single_cert(vec![cert_der], key_der)
53            .unwrap();
54        tls.alpn_protocols = vec![b"h3".to_vec()];
55        tls
56    };
57
58    let quic_server_config =
59        QuicServerConfig::try_from(Arc::new(tls_config)).expect("QuicServerConfig creation failed");
60    let server_config =
61        h3_quinn::quinn::ServerConfig::with_crypto(Arc::new(quic_server_config));
62
63    let endpoint =
64        Endpoint::server(server_config, addr).expect("Failed to bind QUIC endpoint");
65
66    tracing::info!("juiceback QUIC listening on udp://{}", addr);
67
68    let quic_shutdown = shutdown.clone();
69    tokio::select! {
70        biased;
71        _ = quic_shutdown.notified() => {
72            tracing::info!("juiceback QUIC shutting down...");
73            endpoint.wait_idle().await;
74        }
75        _ = run_quic_server(&endpoint, state) => {
76            endpoint.wait_idle().await;
77        }
78    }
79}
80
81async fn run_quic_server(endpoint: &Endpoint, state: Arc<AppState>) {
82    let router = routes::build_router(Arc::clone(&state));
83
84    loop {
85        let conn = match endpoint.accept().await {
86            Some(conn) => conn,
87            None => break,
88        };
89
90        let router = router.clone();
91        tokio::spawn(async move {
92            let conn = match conn.await {
93                Ok(c) => c,
94                Err(e) => {
95                    tracing::warn!("QUIC connection handshake failed: {}", e);
96                    return;
97                }
98            };
99
100            if let Err(e) = handle_h3_conn(H3Connection::new(conn), router).await {
101                tracing::debug!("QUIC connection closed: {}", e);
102            }
103        });
104    }
105}
106
107async fn handle_h3_conn(
108    conn: H3Connection,
109    router: axum::Router,
110) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
111    let mut h3_conn = h3::server::builder().build(conn).await?;
112
113    loop {
114        let resolver = match h3_conn.accept().await {
115            Ok(Some(r)) => r,
116            Ok(None) => break,
117            Err(e) => {
118                tracing::debug!("h3 accept error: {}", e);
119                break;
120            }
121        };
122
123        let mut router = router.clone();
124        tokio::spawn(async move {
125            let (req, mut stream) = match resolver.resolve_request().await {
126                Ok(r) => r,
127                Err(e) => {
128                    tracing::debug!("h3 resolve_request error: {}", e);
129                    return;
130                }
131            };
132
133            if let Err(e) = proxy_axum(&mut router, req, &mut stream).await {
134                tracing::debug!("h3->axum proxy error: {}", e);
135            }
136        });
137    }
138
139    Ok(())
140}
141
142async fn proxy_axum(
143    router: &mut axum::Router,
144    req: Request<()>,
145    stream: &mut h3::server::RequestStream<h3_quinn::BidiStream<Bytes>, Bytes>,
146) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
147    let uri = req.uri().clone();
148    let method = req.method().clone();
149    let headers = req.headers().clone();
150
151    let mut body_bytes = Vec::new();
152    while let Some(chunk) = stream.recv_data().await? {
153        let chunk = chunk.chunk();
154        body_bytes.extend_from_slice(chunk);
155    }
156
157    let mut axum_req = Request::builder()
158        .method(method)
159        .uri(uri)
160        .body(Body::from(body_bytes))
161        .unwrap();
162    *axum_req.headers_mut() = headers;
163
164    let response = Service::call(router, axum_req).await?;
165
166    let status = response.status();
167    let resp_headers = response.headers().clone();
168    let resp_body = response.into_body();
169    let body_bytes: Vec<u8> = resp_body.collect().await?.to_bytes().to_vec();
170
171    let mut builder = Response::builder().status(status);
172    for (k, v) in resp_headers.iter() {
173        builder = builder.header(k, v);
174    }
175    stream
176        .send_response(builder.body(()).unwrap())
177        .await?;
178    stream.send_data(Bytes::from(body_bytes)).await?;
179    stream.finish().await?;
180
181    Ok(())
182}
183
184// -- Client (QUIC push) --
185// Uses quinn + h3 to push file data to juicehost over QUIC/HTTP/3.
186
187#[derive(Debug)]
188struct NoopServerCertVerifier;
189
190impl ServerCertVerifier for NoopServerCertVerifier {
191    fn verify_server_cert(
192        &self,
193        _end_entity: &CertificateDer<'_>,
194        _intermediates: &[CertificateDer<'_>],
195        _server_name: &ServerName<'_>,
196        _ocsp_response: &[u8],
197        _now: UnixTime,
198    ) -> Result<ServerCertVerified, RustlsError> {
199        Ok(ServerCertVerified::assertion())
200    }
201
202    fn verify_tls12_signature(
203        &self,
204        _message: &[u8],
205        _cert: &CertificateDer<'_>,
206        _dss: &DigitallySignedStruct,
207    ) -> Result<HandshakeSignatureValid, RustlsError> {
208        Ok(HandshakeSignatureValid::assertion())
209    }
210
211    fn verify_tls13_signature(
212        &self,
213        _message: &[u8],
214        _cert: &CertificateDer<'_>,
215        _dss: &DigitallySignedStruct,
216    ) -> Result<HandshakeSignatureValid, RustlsError> {
217        Ok(HandshakeSignatureValid::assertion())
218    }
219
220    fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
221        vec![
222            SignatureScheme::RSA_PKCS1_SHA256,
223            SignatureScheme::RSA_PKCS1_SHA384,
224            SignatureScheme::RSA_PKCS1_SHA512,
225            SignatureScheme::ECDSA_NISTP256_SHA256,
226            SignatureScheme::ECDSA_NISTP384_SHA384,
227            SignatureScheme::RSA_PSS_SHA256,
228            SignatureScheme::RSA_PSS_SHA384,
229            SignatureScheme::RSA_PSS_SHA512,
230            SignatureScheme::ED25519,
231            SignatureScheme::ED448,
232        ]
233    }
234}
235
236/// Push file data to juicehost over QUIC/HTTP/3.
237pub async fn push_file_streaming_quic(
238    state: &Arc<AppState>,
239    id: &str,
240    filename: &str,
241    mime_type: &str,
242    chunk_rx: tokio::sync::mpsc::Receiver<Result<Bytes, String>>,
243    host: Option<String>,
244) -> Result<(), String> {
245    let juicehost_url = host.as_deref().unwrap_or(&state.config.juicehost_url);
246    if juicehost_url.is_empty() {
247        return Err("JUICEHOST_URL not set".into());
248    }
249
250    let uri: Uri = juicehost_url
251        .parse()
252        .map_err(|e| format!("invalid juicehost URL: {}", e))?;
253    let hostname = uri.host().ok_or("no hostname in juicehost URL")?.to_string();
254    let quic_port = 6403u16;
255
256    let addr = std::net::SocketAddr::new(
257        hostname
258            .as_str()
259            .parse::<std::net::IpAddr>()
260            .map_err(|e| format!("invalid juicehost IP: {}", e))?,
261        quic_port,
262    );
263
264let _ = rustls::crypto::aws_lc_rs::default_provider()
265    .install_default();
266
267    let mut tls_config = rustls::ClientConfig::builder()
268        .dangerous()
269        .with_custom_certificate_verifier(Arc::new(NoopServerCertVerifier))
270        .with_no_client_auth();
271    tls_config.alpn_protocols = vec![b"h3".to_vec()];
272    tls_config.enable_early_data = true;
273
274    let client_config = quinn::ClientConfig::new(Arc::new(
275        quinn::crypto::rustls::QuicClientConfig::try_from(tls_config)
276            .map_err(|e| format!("QUIC client config: {}", e))?,
277    ));
278
279    let mut endpoint = h3_quinn::quinn::Endpoint::client(
280        std::net::SocketAddr::new(addr.ip(), 0),
281    )
282        .map_err(|e| format!("QUIC endpoint: {}", e))?;
283    endpoint.set_default_client_config(client_config);
284
285    let conn = endpoint
286        .connect(addr, &hostname)
287        .map_err(|e| format!("QUIC connect error: {}", e))?
288        .await
289        .map_err(|e| format!("QUIC handshake failed: {}", e))?;
290
291    let quinn_conn = h3_quinn::Connection::new(conn);
292    let (mut driver, mut send_request) = h3::client::new(quinn_conn)
293        .await
294        .map_err(|e| format!("h3 client setup failed: {}", e))?;
295
296    tokio::spawn(async move {
297        let err = future::poll_fn(|cx| driver.poll_close(cx)).await;
298        if !err.is_h3_no_error() {
299            tracing::warn!("h3 client connection error: {}", err);
300        }
301    });
302
303    let req_uri = format!("{}/internal/file/stream/{}/{}", juicehost_url.trim_end_matches('/'), id, filename);
304    let req = Request::post(&req_uri)
305        .header("x-mime-type", mime_type)
306        .body(())
307        .map_err(|e| format!("failed to build request: {}", e))?;
308
309    let mut stream = send_request
310        .send_request(req)
311        .await
312        .map_err(|e| format!("send_request failed: {}", e))?;
313
314    let mut rx = chunk_rx;
315    let mut has_data = false;
316    while let Some(chunk) = rx.recv().await {
317        has_data = true;
318        let data = chunk.map_err(|e| format!("chunk error: {}", e))?;
319        stream
320            .send_data(data)
321            .await
322            .map_err(|e| format!("send_data failed: {}", e))?;
323    }
324
325    if !has_data {
326        return Err("no data received for QUIC push".into());
327    }
328
329    stream
330        .finish()
331        .await
332        .map_err(|e| format!("finish failed: {}", e))?;
333
334    let resp = stream
335        .recv_response()
336        .await
337        .map_err(|e| format!("recv_response failed: {}", e))?;
338
339    let status = resp.status();
340    let mut body_bytes = Vec::new();
341    while let Some(chunk) = stream
342        .recv_data()
343        .await
344        .map_err(|e| format!("recv_data failed: {}", e))?
345    {
346        body_bytes.extend_from_slice(chunk.chunk());
347    }
348
349    if status.is_success() {
350        tracing::info!("juicehost QUIC push: id={} ok", id);
351        Ok(())
352    } else {
353        let body_text = String::from_utf8_lossy(&body_bytes).to_string();
354        let error_code = serde_json::from_str::<serde_json::Value>(&body_text)
355            .ok()
356            .and_then(|v| v.get("error").and_then(|e| e.as_str().map(|s| s.to_string())))
357            .unwrap_or_default();
358        let detail = serde_json::from_str::<serde_json::Value>(&body_text)
359            .ok()
360            .and_then(|v| v.get("message").and_then(|m| m.as_str().map(|s| s.to_string())))
361            .unwrap_or_else(|| body_text);
362        Err(format!("[{}] {} (status={})", error_code, detail, status.as_u16()))
363    }
364}