karyon_jsonrpc/server/http/
h3.rs1use 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
26type SubSenders = Arc<RwLock<HashMap<SubscriptionID, async_channel::Sender<NewNotification>>>>;
36
37pub(super) async fn accept_h3(server: &Arc<Server>, quic_ep: &QuicEndpoint) {
38 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 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 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
238async 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}