use crate::schema::RequestId;
use crate::{
error::{GenericSendError, TransportError},
message_dispatcher::MessageDispatcher,
utils::CancellationToken,
IoStream,
};
use std::{collections::HashMap, pin::Pin, sync::Arc, time::Duration};
use tokio::task::JoinHandle;
use tokio::{
io::{AsyncBufReadExt, BufReader},
sync::Mutex,
};
pub(crate) const DEFAULT_MESSAGE_CHANNEL_CAPACITY: usize = 36;
pub struct MCPStream {}
impl MCPStream {
#[allow(clippy::too_many_arguments)]
pub fn create<X, R>(
readable: Pin<Box<dyn tokio::io::AsyncRead + Send + Sync>>,
writable: Mutex<Pin<Box<dyn tokio::io::AsyncWrite + Send + Sync>>>,
error_io: IoStream,
pending_requests: Arc<Mutex<HashMap<RequestId, tokio::sync::oneshot::Sender<R>>>>,
request_timeout: Duration,
max_line_length: usize,
cancellation_token: CancellationToken,
channel_capacity: usize,
) -> (
tokio_stream::wrappers::ReceiverStream<X>,
MessageDispatcher<R>,
IoStream,
)
where
R: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
X: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
{
let (tx, rx) = tokio::sync::mpsc::channel::<X>(channel_capacity);
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let reader_token = cancellation_token.clone();
#[allow(clippy::let_underscore_future)]
let _ = Self::spawn_reader(readable, tx, max_line_length, reader_token);
let sender = MessageDispatcher::new(pending_requests, writable, request_timeout);
(stream, sender, error_io)
}
#[allow(clippy::too_many_arguments)]
#[cfg(feature = "streamable-http")]
pub fn create_with_ack<X, R>(
readable: Pin<Box<dyn tokio::io::AsyncRead + Send + Sync>>,
writable: tokio::sync::mpsc::Sender<(
String,
tokio::sync::oneshot::Sender<crate::error::TransportResult<()>>,
)>,
error_io: IoStream,
pending_requests: Arc<Mutex<HashMap<RequestId, tokio::sync::oneshot::Sender<R>>>>,
request_timeout: Duration,
max_line_length: usize,
cancellation_token: CancellationToken,
channel_capacity: usize,
) -> (
tokio_stream::wrappers::ReceiverStream<X>,
MessageDispatcher<R>,
IoStream,
)
where
R: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
X: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
{
let (tx, rx) = tokio::sync::mpsc::channel::<X>(channel_capacity);
let stream = tokio_stream::wrappers::ReceiverStream::new(rx);
let reader_token = cancellation_token.clone();
#[allow(clippy::let_underscore_future)]
let _ = Self::spawn_reader(readable, tx, max_line_length, reader_token);
let sender = MessageDispatcher::new_with_acknowledgement(
pending_requests,
writable,
request_timeout,
);
(stream, sender, error_io)
}
fn spawn_reader<X>(
readable: Pin<Box<dyn tokio::io::AsyncRead + Send + Sync>>,
tx: tokio::sync::mpsc::Sender<X>,
max_line_length: usize,
cancellation_token: CancellationToken,
) -> JoinHandle<Result<(), TransportError>>
where
X: Clone + Send + Sync + serde::de::DeserializeOwned + 'static,
{
tokio::spawn(async move {
let mut reader = BufReader::new(readable);
loop {
tokio::select! {
_ = cancellation_token.cancelled() => {
break;
},
result = read_capped_line(&mut reader, max_line_length) => {
match result {
Ok(LineRead::Eof) => {
break;
}
Ok(LineRead::TooLong) => {
tracing::error!(
"dropping incoming message exceeding {max_line_length} bytes"
);
continue;
}
Ok(LineRead::Line(line)) => {
tracing::trace!("raw payload: {}", &line[..line.len().min(1024)]);
let message: X = match serde_json::from_str(&line) {
Ok(mcp_message) => mcp_message,
Err(_) => {
continue;
}
};
tx.send(message).await.map_err(GenericSendError::new)?;
}
Err(e) => {
tracing::error!("Error reading from readable stream: {e}");
return Err(TransportError::ProcessError(format!(
"Error reading from readable stream: {e}"
)));
}
}
}
}
}
Ok::<(), TransportError>(())
})
}
}
enum LineRead {
Line(String),
TooLong,
Eof,
}
async fn read_capped_line<R>(reader: &mut R, max: usize) -> std::io::Result<LineRead>
where
R: tokio::io::AsyncBufRead + Unpin,
{
let mut buf: Vec<u8> = Vec::new();
loop {
let chunk = reader.fill_buf().await?;
if chunk.is_empty() {
if buf.is_empty() {
return Ok(LineRead::Eof);
}
return Ok(LineRead::Line(line_to_string(buf)));
}
if let Some(pos) = chunk.iter().position(|&b| b == b'\n') {
let consumed = pos + 1;
if buf.len() + pos > max {
reader.consume(consumed);
return Ok(LineRead::TooLong);
}
buf.extend_from_slice(&chunk[..pos]);
reader.consume(consumed);
return Ok(LineRead::Line(line_to_string(buf)));
}
let len = chunk.len();
if buf.len() + len > max {
reader.consume(len);
discard_to_newline(&mut *reader).await?;
return Ok(LineRead::TooLong);
}
buf.extend_from_slice(chunk);
reader.consume(len);
}
}
async fn discard_to_newline<R>(reader: &mut R) -> std::io::Result<()>
where
R: tokio::io::AsyncBufRead + Unpin,
{
loop {
let chunk = reader.fill_buf().await?;
if chunk.is_empty() {
return Ok(());
}
if let Some(pos) = chunk.iter().position(|&b| b == b'\n') {
reader.consume(pos + 1);
return Ok(());
}
let len = chunk.len();
reader.consume(len);
}
}
fn line_to_string(mut buf: Vec<u8>) -> String {
if buf.last() == Some(&b'\r') {
buf.pop();
}
String::from_utf8_lossy(&buf).into_owned()
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::BufReader;
async fn collect_lines(data: &[u8], max: usize) -> Vec<Result<String, &'static str>> {
let mut reader = BufReader::new(data);
let mut out = Vec::new();
loop {
match read_capped_line(&mut reader, max).await.unwrap() {
LineRead::Eof => break,
LineRead::TooLong => out.push(Err("too-long")),
LineRead::Line(line) => out.push(Ok(line)),
}
}
out
}
#[tokio::test]
async fn reads_newline_delimited_lines() {
let out = collect_lines(b"hello\r\nworld\n", 1024).await;
assert_eq!(out, vec![Ok("hello".to_string()), Ok("world".to_string())]);
}
#[tokio::test]
async fn emits_final_line_without_trailing_newline() {
let out = collect_lines(b"tail", 1024).await;
assert_eq!(out, vec![Ok("tail".to_string())]);
}
#[tokio::test]
async fn drops_oversized_line_and_resyncs() {
let mut data = vec![b'a'; 100];
data.push(b'\n');
data.extend_from_slice(b"ok\n");
let out = collect_lines(&data, 10).await;
assert_eq!(out, vec![Err("too-long"), Ok("ok".to_string())]);
}
#[tokio::test]
async fn accepts_line_at_exact_max() {
let data = format!("{}\n", "a".repeat(10));
let out = collect_lines(data.as_bytes(), 10).await;
assert_eq!(out, vec![Ok("a".repeat(10))]);
}
#[tokio::test]
async fn resyncs_after_consecutive_oversized_lines() {
let mut data = vec![b'a'; 100];
data.push(b'\n');
data.extend_from_slice(b"too-big-again\nok\n");
let out = collect_lines(&data, 10).await;
assert_eq!(
out,
vec![Err("too-long"), Err("too-long"), Ok("ok".to_string())]
);
}
#[tokio::test]
async fn reads_empty_line() {
let out = collect_lines(b"\nok\n", 1024).await;
assert_eq!(out, vec![Ok("".to_string()), Ok("ok".to_string())]);
}
#[tokio::test]
async fn drops_line_just_above_max() {
let mut data = vec![b'a'; 11];
data.push(b'\n');
data.extend_from_slice(b"ok\n");
let out = collect_lines(&data, 10).await;
assert_eq!(out, vec![Err("too-long"), Ok("ok".to_string())]);
}
#[tokio::test]
async fn handles_crlf_at_exact_max() {
let data = format!("{}\r\n", "a".repeat(9));
let out = collect_lines(data.as_bytes(), 10).await;
assert_eq!(out, vec![Ok("a".repeat(9))]);
}
}