107 lines
3.4 KiB
Rust
107 lines
3.4 KiB
Rust
use tokio::io::{AsyncBufReadExt, AsyncReadExt, BufReader};
|
|
|
|
/// Reads an MCP (NDJSON or LSP Content-Length prefixed) message from a buffered async reader.
|
|
/// Returns the raw JSON string payload if successful, or None on EOF or error.
|
|
pub async fn read_mcp_message<R: tokio::io::AsyncRead + Unpin>(
|
|
stdin: &mut BufReader<R>,
|
|
) -> Option<String> {
|
|
let mut length = 0;
|
|
let mut line = String::new();
|
|
loop {
|
|
line.clear();
|
|
if stdin.read_line(&mut line).await.unwrap_or(0) == 0 {
|
|
return None;
|
|
}
|
|
|
|
let trimmed_start = line.trim_start();
|
|
if trimmed_start.starts_with('{') {
|
|
let mut val = std::mem::take(&mut line);
|
|
let trimmed_len = val.trim_end().len();
|
|
val.truncate(trimmed_len);
|
|
let start = val.len() - val.trim_start().len();
|
|
if start > 0 {
|
|
val.drain(..start);
|
|
}
|
|
return Some(val);
|
|
}
|
|
|
|
let line = line.trim_end();
|
|
if line.is_empty() {
|
|
break;
|
|
}
|
|
|
|
if line.as_bytes().len() >= 15 && line.as_bytes()[..15].eq_ignore_ascii_case(b"content-length:") {
|
|
length = line[15..].trim().parse().unwrap_or(0);
|
|
}
|
|
}
|
|
|
|
const MAX_MESSAGE_BYTES: usize = 50 * 1024 * 1024; // 50MB safety cap
|
|
|
|
if length == 0 || length > MAX_MESSAGE_BYTES {
|
|
return None;
|
|
}
|
|
|
|
let mut buffer = Vec::with_capacity(length.min(64 * 1024));
|
|
if stdin.take(length as u64).read_to_end(&mut buffer).await.is_err() || buffer.len() != length {
|
|
return None;
|
|
}
|
|
|
|
String::from_utf8(buffer).ok()
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use std::io::Cursor;
|
|
|
|
#[tokio::test]
|
|
async fn test_read_ndjson_message() {
|
|
let input = "{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}\n";
|
|
let mut reader = BufReader::new(Cursor::new(input));
|
|
let msg = read_mcp_message(&mut reader).await;
|
|
assert_eq!(msg, Some("{\"jsonrpc\":\"2.0\",\"id\":1,\"method\":\"ping\"}".to_string()));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_read_content_length_message() {
|
|
let payload = "{\"jsonrpc\":\"2.0\",\"id\":2}";
|
|
let input = format!("Content-Length: {}\r\n\r\n{}", payload.len(), payload);
|
|
let mut reader = BufReader::new(Cursor::new(input));
|
|
let msg = read_mcp_message(&mut reader).await;
|
|
assert_eq!(msg, Some(payload.to_string()));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_read_mcp_message_eof() {
|
|
let input = "";
|
|
let mut reader = BufReader::new(Cursor::new(input));
|
|
let msg = read_mcp_message(&mut reader).await;
|
|
assert_eq!(msg, None);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_read_mcp_message_zero_content_length() {
|
|
let input = "Content-Length: 0\r\n\r\n";
|
|
let mut reader = BufReader::new(Cursor::new(input));
|
|
let msg = read_mcp_message(&mut reader).await;
|
|
assert_eq!(msg, None);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_read_mcp_message_exceeds_max_bytes() {
|
|
let input = "Content-Length: 60000000\r\n\r\n";
|
|
let mut reader = BufReader::new(Cursor::new(input));
|
|
let msg = read_mcp_message(&mut reader).await;
|
|
assert_eq!(msg, None);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_read_mcp_message_read_exact_error() {
|
|
let input = "Content-Length: 1000\r\n\r\nshort";
|
|
let mut reader = BufReader::new(Cursor::new(input));
|
|
let msg = read_mcp_message(&mut reader).await;
|
|
assert_eq!(msg, None);
|
|
}
|
|
}
|
|
|