1use std::pin::pin;
2
3use crate::RawJsonRpcResponse as Response;
5use crate::jsonrpc::{RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame};
6use futures::StreamExt as _;
7use futures::channel::mpsc;
8use serde::Deserialize as _;
9
10enum ParsedIncomingLine {
11 Single(RawJsonRpcMessage),
12 Malformed { raw: String, error: crate::Error },
13 Batch(TransportBatch),
14}
15
16fn parse_incoming_line(line: &str) -> ParsedIncomingLine {
17 let value = match serde_json::from_str::<serde_json::Value>(line) {
18 Ok(value) => value,
19 Err(error) => {
20 tracing::debug!(?error, "Failed to parse incoming JSON-RPC JSON");
21 return ParsedIncomingLine::Malformed {
22 raw: line.to_owned(),
23 error: crate::Error::parse_error().data(serde_json::json!({ "line": line })),
24 };
25 }
26 };
27
28 match value {
29 serde_json::Value::Array(entries) if entries.is_empty() => ParsedIncomingLine::Malformed {
30 raw: line.to_owned(),
31 error: crate::Error::invalid_request(),
32 },
33 serde_json::Value::Array(entries) => {
34 let entries = entries
35 .into_iter()
36 .map(|entry| match RawJsonRpcMessage::deserialize(&entry) {
37 Ok(message) => TransportBatchEntry::message(message),
38 Err(error) => {
39 tracing::debug!(?error, "Invalid JSON-RPC batch entry");
40 TransportBatchEntry::malformed(entry, crate::Error::invalid_request())
41 }
42 })
43 .collect::<Vec<_>>();
44
45 ParsedIncomingLine::Batch(
46 TransportBatch::from_entries(entries)
47 .expect("a parsed non-empty JSON array retains at least one entry"),
48 )
49 }
50 value => match serde_json::from_value(value) {
51 Ok(message) => ParsedIncomingLine::Single(message),
52 Err(error) => {
53 tracing::debug!(?error, "Invalid JSON-RPC message");
54 ParsedIncomingLine::Malformed {
55 raw: line.to_owned(),
56 error: crate::Error::invalid_request(),
57 }
58 }
59 },
60 }
61}
62
63impl TransportFrame {
64 #[must_use]
72 pub fn parse_json(input: &str) -> Self {
73 match parse_incoming_line(input) {
74 ParsedIncomingLine::Single(message) => Self::Single(message),
75 ParsedIncomingLine::Malformed { raw, error } => Self::Malformed { raw, error },
76 ParsedIncomingLine::Batch(batch) => Self::Batch(batch),
77 }
78 }
79
80 pub fn to_json(&self) -> Result<String, crate::Error> {
87 match self {
88 Self::Single(message) => {
89 serde_json::to_string(message).map_err(crate::Error::into_internal_error)
90 }
91 Self::Malformed { raw, .. } => Ok(raw.clone()),
92 Self::Batch(batch) => {
93 serde_json::to_string(batch).map_err(crate::Error::into_internal_error)
94 }
95 }
96 }
97}
98
99async fn transport_outgoing_frames_actor(
113 transport_rx: impl futures::Stream<Item = TransportFrame>,
114 outgoing_lines: impl futures::Sink<String, Error = std::io::Error>,
115) -> Result<(), crate::Error> {
116 use futures::SinkExt;
117 let mut transport_rx = pin!(transport_rx);
118 let mut outgoing_lines = pin!(outgoing_lines);
119
120 while let Some(frame) = transport_rx.next().await {
121 let json_rpc_message = match frame {
122 TransportFrame::Single(message) => message,
123 TransportFrame::Malformed { raw, .. } => {
124 let raw = malformed_line_value(raw)?;
125 tracing::trace!(message = ?raw, "Relaying invalid JSON-RPC value");
126 outgoing_lines
127 .send(raw)
128 .await
129 .map_err(crate::Error::into_internal_error)?;
130 continue;
131 }
132 TransportFrame::Batch(batch) => {
133 let line =
134 serde_json::to_string(&batch).map_err(crate::Error::into_internal_error)?;
135 tracing::trace!(message = %line, "Sending JSON-RPC batch");
136 outgoing_lines
137 .send(line)
138 .await
139 .map_err(crate::Error::into_internal_error)?;
140 continue;
141 }
142 };
143 match serde_json::to_string(&json_rpc_message) {
144 Ok(line) => {
145 tracing::trace!(message = %line, "Sending JSON-RPC message");
146 outgoing_lines
147 .send(line)
148 .await
149 .map_err(crate::Error::into_internal_error)?;
150 }
151
152 Err(serialization_error) => {
153 match json_rpc_message {
154 RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_) => {
155 tracing::error!(
160 ?serialization_error,
161 "Failed to serialize request, ignoring"
162 );
163 }
164 RawJsonRpcMessage::Response(response) => {
165 let id = match response {
168 Response::Result { id, .. } | Response::Error { id, .. } => id,
169 };
170 tracing::error!(
171 ?serialization_error,
172 ?id,
173 "Failed to serialize response, sending internal_error instead"
174 );
175 let error_line = serde_json::to_string(&RawJsonRpcMessage::response(
176 id,
177 Err(crate::Error::internal_error()),
178 ))
179 .unwrap();
180 outgoing_lines
181 .send(error_line)
182 .await
183 .map_err(crate::Error::into_internal_error)?;
184 }
185 }
186 }
187 }
188 }
189 outgoing_lines
190 .close()
191 .await
192 .map_err(crate::Error::into_internal_error)
193}
194
195pub(super) struct LineWriter<W> {
198 writer: std::pin::Pin<Box<W>>,
199 bytes: Vec<u8>,
200 written: usize,
201 closing: bool,
202}
203
204impl<W> LineWriter<W> {
205 pub(super) fn new(writer: W) -> Self {
206 Self {
207 writer: Box::pin(writer),
208 bytes: Vec::new(),
209 written: 0,
210 closing: false,
211 }
212 }
213}
214
215impl<W: futures::AsyncWrite> futures::Sink<String> for LineWriter<W> {
216 type Error = std::io::Error;
217
218 fn poll_ready(
219 self: std::pin::Pin<&mut Self>,
220 cx: &mut std::task::Context<'_>,
221 ) -> std::task::Poll<Result<(), Self::Error>> {
222 self.poll_flush(cx)
223 }
224
225 fn start_send(self: std::pin::Pin<&mut Self>, line: String) -> Result<(), Self::Error> {
226 let this = self.get_mut();
227 this.bytes = line.into_bytes();
228 this.bytes.push(b'\n');
229 this.written = 0;
230 Ok(())
231 }
232
233 fn poll_flush(
234 self: std::pin::Pin<&mut Self>,
235 cx: &mut std::task::Context<'_>,
236 ) -> std::task::Poll<Result<(), Self::Error>> {
237 let this = self.get_mut();
238 while this.written < this.bytes.len() {
239 let count = futures::ready!(
240 this.writer
241 .as_mut()
242 .poll_write(cx, &this.bytes[this.written..])
243 )?;
244 if count == 0 {
245 return std::task::Poll::Ready(Err(std::io::ErrorKind::WriteZero.into()));
246 }
247 this.written += count;
248 }
249 this.bytes.clear();
250 this.written = 0;
251 this.writer.as_mut().poll_flush(cx)
252 }
253
254 fn poll_close(
255 mut self: std::pin::Pin<&mut Self>,
256 cx: &mut std::task::Context<'_>,
257 ) -> std::task::Poll<Result<(), Self::Error>> {
258 if !self.closing {
259 futures::ready!(self.as_mut().poll_flush(cx))?;
260 self.closing = true;
261 }
262 self.get_mut().writer.as_mut().poll_close(cx)
263 }
264}
265
266fn malformed_line_value(raw: String) -> Result<String, crate::Error> {
267 if !raw.contains('\r') && !raw.contains('\n') {
268 return Ok(raw);
269 }
270
271 match serde_json::from_str::<serde_json::Value>(&raw) {
272 Ok(value) => serde_json::to_string(&value),
273 Err(_) => serde_json::to_string(&raw),
274 }
275 .map_err(crate::Error::into_internal_error)
276}
277
278pub(super) async fn transport_outgoing_lines_actor(
279 transport_rx: impl futures::Stream<Item = TransportFrame>,
280 outgoing_lines: impl futures::Sink<String, Error = std::io::Error>,
281) -> Result<(), crate::Error> {
282 transport_outgoing_frames_actor(transport_rx, outgoing_lines).await
283}
284
285pub(super) async fn transport_incoming_lines_actor(
298 incoming_lines: impl futures::Stream<Item = std::io::Result<String>>,
299 transport_tx: mpsc::UnboundedSender<TransportFrame>,
300) -> Result<(), crate::Error> {
301 let mut incoming_lines = pin!(incoming_lines);
302 while let Some(line_result) = incoming_lines.next().await {
303 let line = line_result.map_err(crate::Error::into_internal_error)?;
304 tracing::trace!(message = %line, "Received JSON-RPC message");
305
306 match parse_incoming_line(&line) {
307 ParsedIncomingLine::Single(message) => {
308 transport_tx
309 .unbounded_send(TransportFrame::Single(message))
310 .map_err(crate::Error::into_internal_error)?;
311 }
312 ParsedIncomingLine::Malformed { raw, error } => {
313 transport_tx
314 .unbounded_send(TransportFrame::Malformed { raw, error })
315 .map_err(crate::Error::into_internal_error)?;
316 }
317 ParsedIncomingLine::Batch(entries) => {
318 transport_tx
319 .unbounded_send(TransportFrame::Batch(entries))
320 .map_err(crate::Error::into_internal_error)?;
321 }
322 }
323 }
324 Ok(())
325}
326
327#[cfg(test)]
328mod tests {
329 use std::sync::{Arc, Mutex};
330
331 use super::*;
332 use crate::ErrorCode;
333
334 #[derive(Default)]
335 struct PendingCloseWriter {
336 bytes: Vec<u8>,
337 close_polls: usize,
338 }
339
340 impl futures::AsyncWrite for PendingCloseWriter {
341 fn poll_write(
342 mut self: std::pin::Pin<&mut Self>,
343 _: &mut std::task::Context<'_>,
344 bytes: &[u8],
345 ) -> std::task::Poll<std::io::Result<usize>> {
346 self.bytes.extend_from_slice(bytes);
347 std::task::Poll::Ready(Ok(bytes.len()))
348 }
349
350 fn poll_flush(
351 self: std::pin::Pin<&mut Self>,
352 _: &mut std::task::Context<'_>,
353 ) -> std::task::Poll<std::io::Result<()>> {
354 assert_eq!(self.close_polls, 0, "do not flush again during shutdown");
355 std::task::Poll::Ready(Ok(()))
356 }
357
358 fn poll_close(
359 mut self: std::pin::Pin<&mut Self>,
360 cx: &mut std::task::Context<'_>,
361 ) -> std::task::Poll<std::io::Result<()>> {
362 self.close_polls += 1;
363 if self.close_polls == 1 {
364 cx.waker().wake_by_ref();
365 std::task::Poll::Pending
366 } else {
367 std::task::Poll::Ready(Ok(()))
368 }
369 }
370 }
371
372 #[test]
373 fn byte_writer_flushes_then_continues_pending_shutdown_without_reflushing() {
374 use futures::SinkExt as _;
375 let mut sink = LineWriter::new(PendingCloseWriter::default());
376 futures::executor::block_on(async {
377 sink.send("line".to_string()).await.unwrap();
378 sink.close().await.unwrap();
379 });
380 assert_eq!(sink.writer.bytes, b"line\n");
381 assert_eq!(sink.writer.close_polls, 2);
382 }
383
384 #[test]
385 fn parses_batch_entries_independently() {
386 let ParsedIncomingLine::Batch(batch) = parse_incoming_line(
387 r#"[
388 {"jsonrpc":"2.0","id":1,"method":"one","params":{}},
389 17,
390 {"jsonrpc":"2.0","method":"two","params":{}}
391 ]"#,
392 ) else {
393 panic!("expected a JSON-RPC batch");
394 };
395
396 let entries = batch.iter_results().collect::<Vec<_>>();
397 assert_eq!(entries.len(), 3);
398 assert!(matches!(entries[0], Ok(RawJsonRpcMessage::Request(_))));
399 assert_eq!(entries[1].unwrap_err().code, ErrorCode::InvalidRequest);
400 assert!(matches!(entries[2], Ok(RawJsonRpcMessage::Notification(_))));
401 }
402
403 #[test]
404 fn preserves_every_invalid_member_of_response_batches() {
405 let ParsedIncomingLine::Batch(batch) = parse_incoming_line(
406 r#"[
407 {"jsonrpc":"2.0","id":1,"result":{"ok":true}},
408 17,
409 {"jsonrpc":"2.0","id":2,"result":null,"error":{"code":-32603,"message":"Internal error"}},
410 {"jsonrpc":"2.0","id":3,"error":{"code":-32603,"message":"Internal error"}}
411 ]"#,
412 ) else {
413 panic!("expected a JSON-RPC batch");
414 };
415
416 let entries = batch.iter_results().collect::<Vec<_>>();
417 assert_eq!(entries.len(), 4);
418 assert!(matches!(entries[0], Ok(RawJsonRpcMessage::Response(_))));
419 assert_eq!(entries[1].unwrap_err().code, ErrorCode::InvalidRequest);
420 assert_eq!(entries[2].unwrap_err().code, ErrorCode::InvalidRequest);
421 assert!(matches!(entries[3], Ok(RawJsonRpcMessage::Response(_))));
422 }
423
424 #[test]
425 fn preserves_invalid_value_beside_malformed_response() {
426 let ParsedIncomingLine::Batch(batch) = parse_incoming_line(
427 r#"[
428 17,
429 {"jsonrpc":"2.0","id":1,"result":null,"error":{"code":-32603,"message":"Internal error"}}
430 ]"#,
431 ) else {
432 panic!("expected a JSON-RPC batch");
433 };
434
435 let entries = batch.iter_results().collect::<Vec<_>>();
436 assert_eq!(entries.len(), 2);
437 assert_eq!(entries[0].unwrap_err().code, ErrorCode::InvalidRequest);
438 assert_eq!(entries[1].unwrap_err().code, ErrorCode::InvalidRequest);
439 }
440
441 #[test]
442 fn preserves_entirely_malformed_response_shaped_batch() {
443 let ParsedIncomingLine::Batch(batch) = parse_incoming_line(
444 r#"[
445 {"jsonrpc":"2.0","id":1,"result":null,"error":{"code":-32603,"message":"Internal error"}}
446 ]"#,
447 ) else {
448 panic!("expected a retained JSON-RPC batch");
449 };
450
451 let entries = batch.iter_results().collect::<Vec<_>>();
452 assert_eq!(entries.len(), 1);
453 assert_eq!(entries[0].unwrap_err().code, ErrorCode::InvalidRequest);
454 }
455
456 #[test]
457 fn preserves_invalid_call_shaped_member_beside_response() {
458 let ParsedIncomingLine::Batch(batch) = parse_incoming_line(
459 r#"[
460 {"jsonrpc":"2.0","id":1,"result":null},
461 {"jsonrpc":"2.0","method":1}
462 ]"#,
463 ) else {
464 panic!("expected a JSON-RPC batch");
465 };
466
467 let entries = batch.iter_results().collect::<Vec<_>>();
468 assert_eq!(entries.len(), 2);
469 assert!(matches!(entries[0], Ok(RawJsonRpcMessage::Response(_))));
470 assert_eq!(entries[1].unwrap_err().code, ErrorCode::InvalidRequest);
471 }
472
473 #[test]
474 fn preserves_malformed_response_shaped_member_beside_request() {
475 let ParsedIncomingLine::Batch(batch) = parse_incoming_line(
476 r#"[
477 {"jsonrpc":"2.0","id":1,"method":"one","params":{}},
478 {"jsonrpc":"2.0","id":2,"result":null,"error":{"code":-32603,"message":"Internal error"}}
479 ]"#,
480 ) else {
481 panic!("expected a JSON-RPC batch");
482 };
483
484 let entries = batch.iter_results().collect::<Vec<_>>();
485 assert_eq!(entries.len(), 2);
486 assert!(matches!(entries[0], Ok(RawJsonRpcMessage::Request(_))));
487 assert_eq!(entries[1].unwrap_err().code, ErrorCode::InvalidRequest);
488 }
489
490 #[test]
491 fn preserves_malformed_standalone_response() {
492 let ParsedIncomingLine::Malformed { error, .. } = parse_incoming_line(
493 r#"{"jsonrpc":"2.0","id":1,"result":null,"error":{"code":-32603,"message":"Internal error"}}"#,
494 ) else {
495 panic!("expected one retained invalid response");
496 };
497
498 assert_eq!(error.code, ErrorCode::InvalidRequest);
499 }
500
501 #[test]
502 fn preserves_malformed_call_shaped_standalone_message() {
503 let ParsedIncomingLine::Malformed { error, .. } =
504 parse_incoming_line(r#"{"jsonrpc":"2.0","id":1,"method":"one","result":null}"#)
505 else {
506 panic!("expected one invalid-request error");
507 };
508
509 assert_eq!(error.code, ErrorCode::InvalidRequest);
510 }
511
512 #[test]
513 fn parses_valid_standalone_response() {
514 assert!(matches!(
515 parse_incoming_line(r#"{"jsonrpc":"2.0","id":1,"result":{"ok":true}}"#),
516 ParsedIncomingLine::Single(RawJsonRpcMessage::Response(_))
517 ));
518 }
519
520 #[test]
521 fn all_invalid_batch_defaults_to_call_errors() {
522 let ParsedIncomingLine::Batch(batch) = parse_incoming_line("[1, 2, 3]") else {
523 panic!("expected a JSON-RPC batch");
524 };
525
526 assert_eq!(batch.len(), 3);
527 assert!(
528 batch
529 .iter_results()
530 .all(|entry| entry.unwrap_err().code == ErrorCode::InvalidRequest)
531 );
532 }
533
534 #[test]
535 fn empty_batch_is_an_invalid_request() {
536 let ParsedIncomingLine::Malformed { raw, error } = parse_incoming_line("[]") else {
537 panic!("expected one invalid-request error");
538 };
539
540 assert_eq!(raw, "[]");
541 assert_eq!(error.code, ErrorCode::InvalidRequest);
542 }
543
544 #[test]
545 fn malformed_json_is_a_parse_error() {
546 let ParsedIncomingLine::Malformed { raw, error } = parse_incoming_line("[") else {
547 panic!("expected one parse error");
548 };
549
550 assert_eq!(raw, "[");
551 assert_eq!(error.code, ErrorCode::ParseError);
552 }
553
554 #[test]
555 fn valid_json_with_an_invalid_envelope_is_an_invalid_request() {
556 let ParsedIncomingLine::Malformed { raw, error } = parse_incoming_line("17") else {
557 panic!("expected one invalid-request error");
558 };
559
560 assert_eq!(raw, "17");
561 assert_eq!(error.code, ErrorCode::InvalidRequest);
562 }
563
564 #[tokio::test]
565 async fn multiline_malformed_frame_is_written_as_one_line_value() {
566 let raw = "not json\r\n{\"jsonrpc\":\"2.0\",\"method\":\"injected\"}".to_string();
567 let captured = Arc::new(Mutex::new(Vec::new()));
568 let outgoing = futures::sink::unfold(captured.clone(), |captured, line| async move {
569 captured.lock().unwrap().push(line);
570 Ok::<_, std::io::Error>(captured)
571 });
572
573 transport_outgoing_frames_actor(
574 futures::stream::iter([TransportFrame::Malformed {
575 raw: raw.clone(),
576 error: crate::Error::parse_error(),
577 }]),
578 outgoing,
579 )
580 .await
581 .unwrap();
582
583 let lines = captured.lock().unwrap();
584 assert_eq!(lines.len(), 1);
585 assert!(!lines[0].contains('\r') && !lines[0].contains('\n'));
586 assert_eq!(serde_json::from_str::<String>(&lines[0]).unwrap(), raw);
587 }
588
589 #[test]
590 fn multiline_invalid_json_rpc_value_is_compacted_without_changing_value() {
591 let raw = "{\n \"jsonrpc\": \"2.0\",\n \"method\": 1\n}".to_string();
592 let expected = serde_json::from_str::<serde_json::Value>(&raw).unwrap();
593 let line = malformed_line_value(raw).unwrap();
594
595 assert!(!line.contains('\r') && !line.contains('\n'));
596 assert_eq!(
597 serde_json::from_str::<serde_json::Value>(&line).unwrap(),
598 expected
599 );
600 }
601}