use std::io::BufRead;
use tokio::io::{AsyncBufReadExt, BufReader};
use crate::error::{Error, Result};
pub(crate) enum InputFrame {
Line(String),
Undecodable,
}
pub(crate) fn decode_input_frame(mut raw: Vec<u8>) -> InputFrame {
if raw.last() == Some(&b'\n') {
raw.pop();
if raw.last() == Some(&b'\r') {
raw.pop();
}
}
match String::from_utf8(raw) {
Ok(line) => InputFrame::Line(line),
Err(_) => InputFrame::Undecodable,
}
}
pub(crate) struct FrameReader<R> {
reader: BufReader<R>,
buf: Vec<u8>,
}
impl<R> FrameReader<R>
where
R: tokio::io::AsyncRead + Unpin,
{
pub(crate) fn new(reader: R) -> Self {
Self {
reader: BufReader::new(reader),
buf: Vec::new(),
}
}
pub(crate) async fn next_frame(&mut self) -> Result<Option<InputFrame>> {
let read = self
.reader
.read_until(b'\n', &mut self.buf)
.await
.map_err(|e| Error::Transport(format!("Failed to read input frame: {}", e)))?;
if read == 0 && self.buf.is_empty() {
return Ok(None);
}
Ok(Some(decode_input_frame(std::mem::take(&mut self.buf))))
}
}
pub(crate) fn read_frame_blocking<R: BufRead>(reader: &mut R) -> Result<Option<InputFrame>> {
let mut raw = Vec::new();
let read = reader
.read_until(b'\n', &mut raw)
.map_err(|e| Error::Transport(format!("Failed to read input frame: {}", e)))?;
if read == 0 {
return Ok(None);
}
Ok(Some(decode_input_frame(raw)))
}
pub(crate) fn clean_input_line(line: &str) -> &str {
line.strip_prefix('\u{feff}').unwrap_or(line).trim()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum FrameClass {
Notification,
Response,
Request,
}
pub(crate) fn classify_frame(value: &serde_json::Value) -> FrameClass {
if !value.is_array() && value.get("id").is_none() {
return FrameClass::Notification;
}
if is_response_frame(value) {
return FrameClass::Response;
}
FrameClass::Request
}
pub(crate) fn is_response_frame(value: &serde_json::Value) -> bool {
!value.is_array()
&& value.get("method").is_none()
&& (value.get("result").is_some() || value.get("error").is_some())
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_line(frame: Option<InputFrame>, expected: &str) {
match frame {
Some(InputFrame::Line(line)) => assert_eq!(line, expected),
Some(InputFrame::Undecodable) => panic!("{expected:?} must decode"),
None => panic!("expected a frame, got end of input"),
}
}
fn assert_undecodable(frame: Option<InputFrame>) {
assert!(
matches!(frame, Some(InputFrame::Undecodable)),
"expected an undecodable frame"
);
}
#[test]
fn decoding_strips_the_delimiter_in_both_line_endings() {
assert_line(Some(decode_input_frame(b"{}\n".to_vec())), "{}");
assert_line(Some(decode_input_frame(b"{}\r\n".to_vec())), "{}");
assert_line(Some(decode_input_frame(b"{}".to_vec())), "{}");
}
#[test]
fn decoding_rejects_bytes_rather_than_repairing_them() {
assert_undecodable(Some(decode_input_frame(vec![0xff, 0xfe, b'\n'])));
}
#[tokio::test]
async fn a_bad_frame_costs_only_itself() {
let input: &[u8] = b"\xff\xfe\n{\"id\":1}\n";
let mut frames = FrameReader::new(input);
assert_undecodable(frames.next_frame().await.unwrap());
assert_line(frames.next_frame().await.unwrap(), "{\"id\":1}");
assert!(
frames.next_frame().await.unwrap().is_none(),
"end of input must be reported once the frames are consumed"
);
}
#[test]
fn the_blocking_reader_treats_a_bad_frame_the_same_way() {
let mut input: &[u8] = b"\xff\xfe\n{\"id\":1}\n";
assert_undecodable(read_frame_blocking(&mut input).unwrap());
assert_line(read_frame_blocking(&mut input).unwrap(), "{\"id\":1}");
assert!(read_frame_blocking(&mut input).unwrap().is_none());
}
#[tokio::test]
async fn a_partial_frame_survives_a_cancelled_read() {
let (mut writer, reader) = tokio::io::duplex(256);
let mut frames = FrameReader::new(reader);
let frame = r#"{"jsonrpc":"2.0","id":2,"result":{"tools":[]}}"#;
tokio::io::AsyncWriteExt::write_all(&mut writer, &frame.as_bytes()[..10])
.await
.unwrap();
assert!(
tokio::time::timeout(std::time::Duration::from_millis(10), frames.next_frame())
.await
.is_err(),
"a partial frame must remain pending until its newline arrives"
);
tokio::io::AsyncWriteExt::write_all(&mut writer, &frame.as_bytes()[10..])
.await
.unwrap();
tokio::io::AsyncWriteExt::write_all(&mut writer, b"\n")
.await
.unwrap();
assert_line(frames.next_frame().await.unwrap(), frame);
}
#[test]
fn test_clean_input_line_no_bom() {
assert_eq!(
clean_input_line(r#"{"jsonrpc":"2.0"}"#),
r#"{"jsonrpc":"2.0"}"#
);
}
#[test]
fn test_clean_input_line_strips_leading_bom() {
let with_bom = "\u{feff}{\"jsonrpc\":\"2.0\"}";
assert_eq!(clean_input_line(with_bom), r#"{"jsonrpc":"2.0"}"#);
}
#[test]
fn test_clean_input_line_strips_bom_then_trims() {
let input = "\u{feff} {\"id\":1}\n";
assert_eq!(clean_input_line(input), r#"{"id":1}"#);
}
#[test]
fn test_clean_input_line_does_not_strip_internal_bom() {
let input = "{\"text\":\"hi\u{feff}there\"}";
assert_eq!(clean_input_line(input), input);
}
#[test]
fn test_clean_input_line_empty() {
assert_eq!(clean_input_line(""), "");
assert_eq!(clean_input_line("\u{feff}"), "");
assert_eq!(clean_input_line(" \n\t"), "");
}
#[test]
fn classify_frame_notification_has_no_id() {
let value = serde_json::json!({"jsonrpc": "2.0", "method": "notifications/initialized"});
assert_eq!(classify_frame(&value), FrameClass::Notification);
}
#[test]
fn classify_frame_response_has_id_no_method_and_result() {
let value = serde_json::json!({"jsonrpc": "2.0", "id": 1, "result": {}});
assert_eq!(classify_frame(&value), FrameClass::Response);
}
#[test]
fn classify_frame_response_has_id_no_method_and_error() {
let value =
serde_json::json!({"jsonrpc": "2.0", "id": 1, "error": {"code": -1, "message": "x"}});
assert_eq!(classify_frame(&value), FrameClass::Response);
}
#[test]
fn classify_frame_request_has_id_and_method() {
let value = serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "tools/list"});
assert_eq!(classify_frame(&value), FrameClass::Request);
}
#[test]
fn classify_frame_batch_array_is_always_request() {
let value = serde_json::json!([
{"jsonrpc": "2.0", "id": 1, "method": "a"},
{"jsonrpc": "2.0", "method": "b"},
]);
assert_eq!(classify_frame(&value), FrameClass::Request);
}
#[test]
fn classify_frame_id_present_but_null_is_not_a_notification() {
let value = serde_json::json!({"jsonrpc": "2.0", "id": null, "method": "tools/list"});
assert_eq!(classify_frame(&value), FrameClass::Request);
}
#[test]
fn classify_frame_id_present_but_null_can_still_be_a_response() {
let value = serde_json::json!({"jsonrpc": "2.0", "id": null, "result": {}});
assert_eq!(classify_frame(&value), FrameClass::Response);
}
#[test]
fn classify_frame_id_no_method_no_result_or_error_is_a_request_not_a_response() {
let value = serde_json::json!({"jsonrpc": "2.0", "id": 1});
assert_eq!(classify_frame(&value), FrameClass::Request);
}
#[test]
fn is_response_frame_matches_the_response_shape() {
assert!(is_response_frame(
&serde_json::json!({"jsonrpc": "2.0", "id": 1, "result": {}})
));
assert!(is_response_frame(
&serde_json::json!({"jsonrpc": "2.0", "id": 1, "error": {"code": -1}})
));
}
#[test]
fn is_response_frame_rejects_a_request_shape() {
assert!(!is_response_frame(
&serde_json::json!({"jsonrpc": "2.0", "id": 1, "method": "tools/list"})
));
assert!(!is_response_frame(
&serde_json::json!({"jsonrpc": "2.0", "id": 1})
));
}
#[test]
fn is_response_frame_does_not_itself_require_an_id() {
assert!(is_response_frame(&serde_json::json!({"result": {}})));
}
#[test]
fn is_response_frame_rejects_a_batch_array() {
assert!(!is_response_frame(&serde_json::json!([
{"result": {}},
])));
}
}