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

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