Files
mcp-memory/mcp-stdio/src/lib.rs
T

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);
}
}