121 lines
3.7 KiB
Rust
121 lines
3.7 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 = line.trim();
|
|
if trimmed.starts_with('{') || trimmed.starts_with('[') {
|
|
return Some(trimmed.to_string());
|
|
}
|
|
|
|
let trimmed_line = line.trim_end();
|
|
if trimmed_line.is_empty() {
|
|
break;
|
|
}
|
|
|
|
if trimmed_line.len() >= 15
|
|
&& trimmed_line.as_bytes()[..15].eq_ignore_ascii_case(b"content-length:")
|
|
{
|
|
length = trimmed_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_batch_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_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);
|
|
}
|
|
}
|