1use 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
29pub 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}