Skip to main content

karyon_jsonrpc/server/http/
h3.rs

1//! HTTP/3 over QUIC via h3 / h3-quinn.
2
3use std::{collections::HashMap, sync::Arc};
4
5use bytes::{Buf, Bytes};
6use h3_quinn::Connection;
7use hyper::{Method, Request, Response, StatusCode};
8use log::{debug, error};
9
10use karyon_core::async_runtime::lock::RwLock;
11
12use karyon_net::quic::{QuicConn, QuicEndpoint, QuicIncoming};
13
14use crate::{
15    error::{Error, Result},
16    message::{self, SubscriptionID},
17    server::{
18        channel::{Channel, NewNotification},
19        http::{ERR_BODY_TOO_LARGE, ERR_METHOD_NOT_ALLOWED, MAX_HTTP_BODY_SIZE},
20        Server, CHANNEL_SUBSCRIPTION_BUFFER_SIZE, FAILED_TO_PARSE_ERROR_MSG,
21    },
22};
23
24type H3Stream = h3::server::RequestStream<h3_quinn::BidiStream<Bytes>, Bytes>;
25
26/// Map from `SubscriptionID` to the channel sender that feeds that
27/// subscription's HTTP/3 reply stream. Used by `dispatch_subs_task`
28/// to look up the right sender for each incoming `NewNotification`
29/// and forward the notification to its stream.
30///
31/// Inserted on subscribe, removed on unsubscribe (or when the stream
32/// ends). `RwLock` because the dispatcher reads it on every
33/// notification (frequent), while subscribe / unsubscribe write
34/// rarely.
35type SubSenders = Arc<RwLock<HashMap<SubscriptionID, async_channel::Sender<NewNotification>>>>;
36
37pub(super) async fn accept_h3(server: &Arc<Server>, quic_ep: &QuicEndpoint) {
38    // Only the accept runs here; the handshake runs in the spawned task.
39    let incoming = match quic_ep.accept_incoming().await {
40        Ok(incoming) => incoming,
41        Err(err) => {
42            error!("Accept QUIC for HTTP/3: {err}");
43            return;
44        }
45    };
46
47    server
48        .task_group
49        .spawn(serve_incoming_task(server.clone(), incoming));
50}
51
52async fn serve_incoming_task(server: Arc<Server>, incoming: QuicIncoming) -> Result<()> {
53    let peer = incoming.peer_endpoint();
54    match incoming.handshake().await {
55        Ok(quic_conn) => serve_conn_task(server, quic_conn).await,
56        Err(err) => {
57            debug!("QUIC handshake with {peer} failed: {err}");
58            Ok(())
59        }
60    }
61}
62
63async fn serve_conn_task(server: Arc<Server>, quic_conn: QuicConn) -> Result<()> {
64    let peer = quic_conn
65        .peer_endpoint()
66        .map(|e| e.to_string())
67        .unwrap_or_default();
68    if let Err(err) = serve_conn(server, quic_conn).await {
69        debug!("HTTP/3 from {peer} closed: {err}");
70    }
71    Ok(())
72}
73
74async fn serve_conn(server: Arc<Server>, quic_conn: QuicConn) -> Result<()> {
75    let h3_conn = Connection::new(quic_conn.inner().clone());
76    let mut h3_server: h3::server::Connection<Connection, Bytes> =
77        h3::server::Connection::new(h3_conn)
78            .await
79            .map_err(|e| Error::HttpError(format!("H3 handshake: {e}")))?;
80
81    // Per-connection channel feeds all pubsub output.
82    let (ch_tx, ch_rx) = async_channel::bounded(CHANNEL_SUBSCRIPTION_BUFFER_SIZE);
83    let channel = Channel::new(ch_tx);
84
85    let sub_senders: SubSenders = Arc::new(RwLock::new(HashMap::new()));
86
87    // Dispatcher task: route notifications from the connection-wide
88    // channel to the right per-subscription sender.
89    server
90        .task_group
91        .spawn(dispatch_subs_task(ch_rx, sub_senders.clone()));
92
93    loop {
94        match h3_server.accept().await {
95            Ok(Some(resolver)) => {
96                let (req, stream) = match resolver.resolve_request().await {
97                    Ok(v) => v,
98                    Err(err) => {
99                        error!("Resolve HTTP/3 request: {err}");
100                        continue;
101                    }
102                };
103                server.task_group.spawn(handle_request_task(
104                    server.clone(),
105                    channel.clone(),
106                    sub_senders.clone(),
107                    req,
108                    stream,
109                ));
110            }
111            Ok(None) => break,
112            Err(err) => {
113                debug!("HTTP/3 connection error: {err}");
114                break;
115            }
116        }
117    }
118
119    for (_, sender) in sub_senders.write().await.drain() {
120        sender.close();
121    }
122    channel.close();
123    Ok(())
124}
125
126async fn dispatch_subs_task(
127    ch_rx: async_channel::Receiver<NewNotification>,
128    sub_senders: SubSenders,
129) -> Result<()> {
130    while let Ok(nt) = ch_rx.recv().await {
131        let subs = sub_senders.read().await;
132        if let Some(sender) = subs.get(&nt.sub_id) {
133            let _ = sender.send(nt).await;
134        }
135    }
136    Ok(())
137}
138
139async fn handle_request_task(
140    server: Arc<Server>,
141    channel: Arc<Channel>,
142    sub_senders: SubSenders,
143    req: Request<()>,
144    stream: H3Stream,
145) -> Result<()> {
146    if let Err(err) = handle_h3_request(server, channel, sub_senders, req, stream).await {
147        error!("Handle HTTP/3 request: {err}");
148    }
149    Ok(())
150}
151
152async fn handle_h3_request(
153    server: Arc<Server>,
154    channel: Arc<Channel>,
155    sub_senders: SubSenders,
156    req: Request<()>,
157    mut stream: H3Stream,
158) -> Result<()> {
159    if req.method() != Method::POST {
160        h3_send(
161            &mut stream,
162            StatusCode::METHOD_NOT_ALLOWED,
163            ERR_METHOD_NOT_ALLOWED.as_bytes(),
164        )
165        .await?;
166        return Ok(());
167    }
168
169    let mut body = Vec::new();
170    while let Some(chunk) = stream.recv_data().await.map_err(h3_err)? {
171        body.extend_from_slice(Buf::chunk(&chunk));
172        if body.len() as u64 > MAX_HTTP_BODY_SIZE {
173            h3_send(
174                &mut stream,
175                StatusCode::PAYLOAD_TOO_LARGE,
176                ERR_BODY_TOO_LARGE.as_bytes(),
177            )
178            .await?;
179            return Ok(());
180        }
181    }
182
183    let msg: serde_json::Value = match serde_json::from_slice(&body) {
184        Ok(v) => v,
185        Err(_) => {
186            let resp = message::Response {
187                error: Some(message::Error {
188                    code: message::PARSE_ERROR_CODE,
189                    message: FAILED_TO_PARSE_ERROR_MSG.to_string(),
190                    data: None,
191                }),
192                ..Default::default()
193            };
194            h3_send(
195                &mut stream,
196                StatusCode::OK,
197                &serde_json::to_vec(&resp).unwrap(),
198            )
199            .await?;
200            return Ok(());
201        }
202    };
203
204    let is_subscribe = msg
205        .get("method")
206        .and_then(|m| m.as_str())
207        .map(|m| m.ends_with("_subscribe") && !m.ends_with("_unsubscribe"))
208        .unwrap_or(false);
209
210    if is_subscribe {
211        return handle_h3_subscribe(server, sub_senders, msg, stream).await;
212    }
213
214    let response = server.handle_request(Some(channel.clone()), msg).await;
215    debug!("--> {response}");
216
217    if response.error.is_none() {
218        if let Ok(rpc_req) = serde_json::from_slice::<message::Request>(&body) {
219            if let Some(params) = &rpc_req.params {
220                if let Ok(sub_id) = serde_json::from_value::<SubscriptionID>(params.clone()) {
221                    if let Some(sender) = sub_senders.write().await.remove(&sub_id) {
222                        sender.close();
223                    }
224                }
225            }
226        }
227    }
228
229    h3_send(
230        &mut stream,
231        StatusCode::OK,
232        &serde_json::to_vec(&response).unwrap(),
233    )
234    .await?;
235    Ok(())
236}
237
238/// Subscribe: respond, then stream notifications on the same stream.
239async fn handle_h3_subscribe(
240    server: Arc<Server>,
241    sub_senders: SubSenders,
242    msg: serde_json::Value,
243    mut stream: H3Stream,
244) -> Result<()> {
245    let (sub_tx, sub_rx) = async_channel::bounded(CHANNEL_SUBSCRIPTION_BUFFER_SIZE);
246    let sub_channel = Channel::new(sub_tx.clone());
247
248    let response = server.handle_request(Some(sub_channel.clone()), msg).await;
249
250    if response.error.is_some() {
251        h3_send(
252            &mut stream,
253            StatusCode::OK,
254            &serde_json::to_vec(&response).unwrap(),
255        )
256        .await?;
257        return Ok(());
258    }
259
260    let sub_id = response
261        .result
262        .as_ref()
263        .and_then(|v| serde_json::from_value::<SubscriptionID>(v.clone()).ok())
264        .ok_or_else(|| Error::InvalidMsg("Missing subscription id".into()))?;
265
266    sub_senders.write().await.insert(sub_id, sub_tx);
267
268    let resp = Response::builder()
269        .status(StatusCode::OK)
270        .header("Content-Type", "application/json")
271        .body(())
272        .unwrap();
273    stream.send_response(resp).await.map_err(h3_err)?;
274
275    debug!("--> {response}");
276    let json = serde_json::to_vec(&response).unwrap();
277    stream.send_data(Bytes::from(json)).await.map_err(h3_err)?;
278
279    let encoder = server.config.notification_encoder;
280
281    while let Ok(nt) = sub_rx.recv().await {
282        let notification = encoder(nt);
283        debug!("--> {notification}");
284        let json = serde_json::to_vec(&serde_json::json!(notification)).unwrap();
285        if stream.send_data(Bytes::from(json)).await.is_err() {
286            break;
287        }
288    }
289
290    sub_senders.write().await.remove(&sub_id);
291    sub_channel.close();
292    let _ = stream.finish().await;
293    Ok(())
294}
295
296async fn h3_send(stream: &mut H3Stream, status: StatusCode, body: &[u8]) -> Result<()> {
297    let resp = Response::builder()
298        .status(status)
299        .header("Content-Type", "application/json")
300        .body(())
301        .map_err(|e| Error::HttpError(e.to_string()))?;
302    stream.send_response(resp).await.map_err(h3_err)?;
303    stream
304        .send_data(Bytes::from(body.to_vec()))
305        .await
306        .map_err(h3_err)?;
307    stream.finish().await.map_err(h3_err)?;
308    Ok(())
309}
310
311fn h3_err(e: impl std::fmt::Display) -> Error {
312    Error::HttpError(e.to_string())
313}