1use 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
30fn 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#[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
236pub 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}