use crate::mcp::{McpError, McpResult};
use std::io::BufRead;
pub(crate) const DEFAULT_MAX_SSE_EVENT_BYTES: usize = 1_048_576;
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct SseEvent {
pub(crate) id: Option<String>,
pub(crate) event_type: Option<String>,
pub(crate) retry: Option<u64>,
pub(crate) data: String,
}
#[derive(Debug, Default)]
pub(crate) struct SseParser {
data_lines: Vec<String>,
id: Option<String>,
event_type: Option<String>,
retry: Option<u64>,
}
impl SseParser {
pub(crate) fn push_line(&mut self, line: &str) -> Option<SseEvent> {
let line = line.strip_suffix('\r').unwrap_or(line);
if line.is_empty() {
return self.dispatch();
}
if line.starts_with(':') {
return None;
}
let (field, value) = match line.split_once(':') {
Some((field, value)) => (field, value.strip_prefix(' ').unwrap_or(value)),
None => (line, ""),
};
match field {
"data" => self.data_lines.push(value.to_string()),
"id" => self.id = Some(value.to_string()),
"event" => self.event_type = Some(value.to_string()),
"retry" => {
if let Ok(retry) = value.parse::<u64>() {
self.retry = Some(retry);
}
}
_ => {}
}
None
}
fn dispatch(&mut self) -> Option<SseEvent> {
if self.data_lines.is_empty()
&& self.id.is_none()
&& self.event_type.is_none()
&& self.retry.is_none()
{
return None;
}
Some(SseEvent {
id: self.id.take(),
event_type: self.event_type.take(),
retry: self.retry.take(),
data: std::mem::take(&mut self.data_lines).join("\n"),
})
}
}
const SSE_BUFFER_COMPACTION_THRESHOLD: usize = 8 * 1024;
#[derive(Debug, Default)]
pub(crate) struct SseDecoder {
buffer: Vec<u8>,
cursor: usize,
parser: SseParser,
total: usize,
}
impl SseDecoder {
pub(crate) fn push_chunk(
&mut self,
chunk: &[u8],
max_bytes: usize,
) -> McpResult<Vec<SseEvent>> {
self.buffer.extend_from_slice(chunk);
let mut events = Vec::new();
while let Some(relative_newline) = self.buffer[self.cursor..]
.iter()
.position(|byte| *byte == b'\n')
{
let line_start = self.cursor;
let newline = line_start + relative_newline;
self.cursor = newline + 1;
self.total = self
.total
.saturating_add(newline.saturating_sub(line_start) + 1);
if self.total > max_bytes {
return Err(McpError::Transport(format!(
"MCP SSE event exceeded {max_bytes} bytes"
)));
}
let line = std::str::from_utf8(&self.buffer[line_start..newline])
.map_err(McpError::transport)?;
if let Some(event) = self.parser.push_line(line) {
self.total = 0;
events.push(event);
}
}
self.compact_buffer();
if self.pending_len().saturating_add(self.total) > max_bytes {
return Err(McpError::Transport(format!(
"MCP SSE event exceeded {max_bytes} bytes"
)));
}
Ok(events)
}
pub(crate) fn finish(&mut self, max_bytes: usize) -> McpResult<Option<SseEvent>> {
if self.pending_len() == 0 {
return Ok(self.parser.dispatch());
}
if self.pending_len().saturating_add(self.total) > max_bytes {
return Err(McpError::Transport(format!(
"MCP SSE event exceeded {max_bytes} bytes"
)));
}
let line = std::str::from_utf8(&self.buffer[self.cursor..]).map_err(McpError::transport)?;
let event = self
.parser
.push_line(line)
.or_else(|| self.parser.dispatch());
self.buffer.clear();
self.cursor = 0;
Ok(event)
}
fn pending_len(&self) -> usize {
self.buffer.len().saturating_sub(self.cursor)
}
fn compact_buffer(&mut self) {
if self.cursor == 0 {
return;
}
if self.cursor == self.buffer.len() {
self.buffer.clear();
self.cursor = 0;
} else if self.cursor >= SSE_BUFFER_COMPACTION_THRESHOLD
&& self.cursor >= self.buffer.len() / 2
{
self.buffer.drain(..self.cursor);
self.cursor = 0;
}
}
}
pub(crate) fn read_sse_event(
reader: &mut impl BufRead,
max_bytes: usize,
) -> McpResult<Option<SseEvent>> {
let mut parser = SseParser::default();
let mut line = Vec::new();
let mut total = 0usize;
loop {
if !read_sse_line(reader, &mut line, &mut total, max_bytes)? {
return Ok(parser.dispatch());
}
while line.last() == Some(&b'\r') {
line.pop();
}
let line = std::str::from_utf8(&line).map_err(McpError::transport)?;
if let Some(event) = parser.push_line(line) {
return Ok(Some(event));
}
}
}
fn read_sse_line(
reader: &mut impl BufRead,
line: &mut Vec<u8>,
total: &mut usize,
max_bytes: usize,
) -> McpResult<bool> {
line.clear();
loop {
let available = reader.fill_buf().map_err(McpError::transport)?;
if available.is_empty() {
return Ok(!line.is_empty());
}
let mut consumed = 0;
let mut complete = false;
for byte in available {
if *total == max_bytes {
break;
}
*total += 1;
consumed += 1;
if *byte == b'\n' {
complete = true;
break;
}
line.push(*byte);
}
reader.consume(consumed);
if complete {
return Ok(true);
}
if *total == max_bytes {
if reader.fill_buf().map_err(McpError::transport)?.is_empty() {
return Ok(!line.is_empty());
}
return Err(McpError::Transport(format!(
"MCP SSE event exceeded {max_bytes} bytes"
)));
}
}
}
pub(crate) fn jsonrpc_message_from_event(
event: &SseEvent,
) -> McpResult<Option<crate::mcp::jsonrpc::JsonRpcMessage>> {
if event.data.trim().is_empty() {
return Ok(None);
}
serde_json::from_str(&event.data)
.map(Some)
.map_err(|error| McpError::Transport(format!("MCP SSE JSON parse failed: {error}")))
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn parses_multi_line_data_and_boundary() {
let mut parser = SseParser::default();
assert!(parser.push_line("data: {\"a\":1").is_none());
assert!(parser.push_line("data: }").is_none());
let event = parser.push_line("").unwrap();
assert_eq!(event.data, "{\"a\":1\n}");
}
#[test]
fn ignores_comments_and_parses_id_event_retry() {
let mut reader = Cursor::new(b":hello\nid: 7\nevent: message\nretry: 1500\ndata: {}\n\n");
let event = read_sse_event(&mut reader, DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap()
.unwrap();
assert_eq!(event.id.as_deref(), Some("7"));
assert_eq!(event.event_type.as_deref(), Some("message"));
assert_eq!(event.retry, Some(1500));
assert_eq!(event.data, "{}");
}
#[test]
fn dispatches_priming_event_without_data() {
let mut reader = Cursor::new(b"id: boot\nretry: 100\n\n");
let event = read_sse_event(&mut reader, DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap()
.unwrap();
assert_eq!(event.id.as_deref(), Some("boot"));
assert_eq!(event.retry, Some(100));
assert!(event.data.is_empty());
assert!(jsonrpc_message_from_event(&event).unwrap().is_none());
}
#[test]
fn ignores_malformed_unknown_lines_and_empty_events() {
let mut parser = SseParser::default();
assert!(parser.push_line("wat").is_none());
assert!(parser.push_line("").is_none());
assert!(parser.push_line("").is_none());
}
#[test]
fn rejects_oversized_events() {
let mut reader = Cursor::new(b"data: 12345\n\n");
let error = read_sse_event(&mut reader, 4).unwrap_err().to_string();
assert!(error.contains("exceeded"), "{error}");
}
#[test]
fn rejects_newline_free_line_at_event_byte_limit() {
let input = vec![b'x'; 9];
let mut reader = Cursor::new(input);
let error = read_sse_event(&mut reader, 8).unwrap_err().to_string();
assert!(error.contains("MCP SSE event exceeded 8 bytes"), "{error}");
assert_eq!(reader.position(), 8);
}
#[test]
fn accepts_unterminated_event_at_exact_byte_limit() {
let input = b"data: x";
let mut reader = Cursor::new(input);
let event = read_sse_event(&mut reader, input.len()).unwrap().unwrap();
assert_eq!(event.data, "x");
}
#[test]
fn parses_jsonrpc_message_from_event() {
let event = SseEvent {
id: None,
event_type: None,
retry: None,
data: r#"{"jsonrpc":"2.0","id":1,"result":{}}"#.to_string(),
};
assert!(jsonrpc_message_from_event(&event).unwrap().is_some());
}
#[test]
fn decoder_handles_split_utf8_crlf_and_event_boundaries() {
let mut decoder = SseDecoder::default();
assert!(
decoder
.push_chunk(b"data: caf", DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap()
.is_empty()
);
assert!(
decoder
.push_chunk("Ă©\r\n".as_bytes(), DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap()
.is_empty()
);
let events = decoder
.push_chunk(b"data: second\r\n\r\n", DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].data, "café\nsecond");
}
#[test]
fn finish_dispatches_valid_unterminated_event() {
let mut decoder = SseDecoder::default();
decoder
.push_chunk(b"data: final", DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap();
assert_eq!(
decoder
.finish(DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap()
.unwrap()
.data,
"final"
);
}
#[test]
fn finish_rejects_invalid_trailing_utf8() {
let mut decoder = SseDecoder::default();
decoder
.push_chunk(b"data: \xff", DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap();
let error = decoder
.finish(DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap_err()
.to_string();
assert!(
error.contains("UTF-8") || error.contains("utf-8"),
"{error}"
);
}
#[test]
fn finish_rejects_oversized_buffered_event() {
let mut decoder = SseDecoder::default();
decoder.push_chunk(b"data: too-long", 8).unwrap_err();
let mut decoder = SseDecoder::default();
decoder
.push_chunk(b"data: x", DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap();
let error = decoder.finish(4).unwrap_err().to_string();
assert!(error.contains("exceeded"), "{error}");
}
#[test]
fn decoder_handles_many_short_lines_without_per_line_buffer_drain() {
const LINE_COUNT: usize = 32_768;
let mut input = Vec::with_capacity(LINE_COUNT * b"data: x\n".len() + 1);
for _ in 0..LINE_COUNT {
input.extend_from_slice(b"data: x\n");
}
input.push(b'\n');
let mut decoder = SseDecoder::default();
let events = decoder
.push_chunk(&input, DEFAULT_MAX_SSE_EVENT_BYTES)
.unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].data.split('\n').count(), LINE_COUNT);
assert_eq!(events[0].data.len(), LINE_COUNT * 2 - 1);
}
const PARTITION_FIXTURE: &[u8] =
b"id: 7\nevent: message\nretry: 1500\ndata: caf\xc3\xa9\ndata: second\r\n\r\ndata: final without newline";
proptest::proptest! {
#[test]
fn decoder_matches_one_shot_oracle_for_arbitrary_byte_partitions(
partition_points in proptest::collection::vec(
proptest::bool::ANY,
PARTITION_FIXTURE.len()..=PARTITION_FIXTURE.len(),
)
) {
let input = PARTITION_FIXTURE;
let mut oracle = SseDecoder::default();
let mut expected = oracle.push_chunk(input, DEFAULT_MAX_SSE_EVENT_BYTES).unwrap();
if let Some(event) = oracle.finish(DEFAULT_MAX_SSE_EVENT_BYTES).unwrap() {
expected.push(event);
}
let mut decoder = SseDecoder::default();
let mut actual = Vec::new();
let mut start = 0;
for index in 0..input.len() {
if partition_points[index] {
actual.extend(decoder.push_chunk(&input[start..=index], DEFAULT_MAX_SSE_EVENT_BYTES).unwrap());
start = index + 1;
}
}
actual.extend(decoder.push_chunk(&input[start..], DEFAULT_MAX_SSE_EVENT_BYTES).unwrap());
if let Some(event) = decoder.finish(DEFAULT_MAX_SSE_EVENT_BYTES).unwrap() {
actual.push(event);
}
assert_eq!(actual, expected);
}
}
}