Skip to main content

agent_client_protocol/jsonrpc/
transport_actor.rs

1use std::pin::pin;
2
3// Types re-exported from crate root
4use 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    /// Parse one JSON-RPC wire value while preserving batch boundaries.
65    ///
66    /// Every malformed value is retained as an explicit malformed frame or
67    /// batch entry. Standalone malformed input keeps its original text; batch
68    /// entries keep their parsed JSON values, source order, and batch boundary,
69    /// though reserialization may normalize whitespace. Protocol actors decide
70    /// whether a malformed value is call-shaped and requires an Error Response.
71    #[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    /// Serialize this frame to its JSON-RPC wire representation.
81    ///
82    /// # Errors
83    ///
84    /// Returns an internal error if a valid message or batch cannot be
85    /// serialized. Malformed frames return their original wire text unchanged.
86    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
99/// Transport outgoing actor for line streams: serializes [`TransportFrame`] values and yields
100/// lines.
101///
102/// This is a line-based variant of `transport_outgoing_actor` that works with a Sink<String>
103/// instead of an AsyncWrite byte stream. This enables interception of lines before they are
104/// written to the underlying transport.
105///
106/// This actor handles transport mechanics:
107/// - Serializes single messages, malformed values, and batches to JSON strings
108/// - Yields newline-terminated strings
109/// - Handles serialization errors
110///
111/// This is the transport layer - it has no knowledge of protocol semantics (IDs, correlation, etc.).
112async 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                        // If we failed to serialize a request,
156                        // just ignore it.
157                        //
158                        // Q: (Maybe it'd be nice to "reply" with an error?)
159                        tracing::error!(
160                            ?serialization_error,
161                            "Failed to serialize request, ignoring"
162                        );
163                    }
164                    RawJsonRpcMessage::Response(response) => {
165                        // If we failed to serialize a *response*,
166                        // send an error in response.
167                        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
195/// A newline writer whose sink close reaches the actual write half. An unfold
196/// sink can flush each line, but has no way to forward `AsyncWrite::close`.
197pub(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
285/// Transport incoming actor for line streams: parses lines into [`TransportFrame`] values.
286///
287/// This is a line-based variant of `transport_incoming_actor` that works with a
288/// Stream<Item = io::Result<String>> instead of an AsyncRead byte stream. This enables
289/// interception of lines before they are parsed.
290///
291/// This actor handles transport mechanics:
292/// - Reads lines from the stream
293/// - Parses individual messages and retains batch arrays in entry order
294/// - Handles malformed JSON, empty batches, and invalid batch entries
295///
296/// This is the transport layer - it has no knowledge of protocol semantics.
297pub(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}