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"),
})
}
}
#[derive(Debug, Default)]
pub(crate) struct SseDecoder {
buffer: Vec<u8>,
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(newline) = self.buffer.iter().position(|byte| *byte == b'\n') {
let line_bytes: Vec<u8> = self.buffer.drain(..=newline).collect();
let line_bytes = &line_bytes[..line_bytes.len() - 1];
self.total = self.total.saturating_add(line_bytes.len() + 1);
if self.total > max_bytes {
return Err(McpError::Transport(format!(
"MCP SSE event exceeded {max_bytes} bytes"
)));
}
let line = std::str::from_utf8(line_bytes).map_err(McpError::transport)?;
if let Some(event) = self.parser.push_line(line) {
self.total = 0;
events.push(event);
}
}
if self.buffer.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.buffer.is_empty() {
return Ok(self.parser.dispatch());
}
if self.buffer.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)
.map_err(McpError::transport)?
.to_string();
self.buffer.clear();
Ok(self
.parser
.push_line(&line)
.or_else(|| self.parser.dispatch()))
}
}
pub(crate) fn read_sse_event(
reader: &mut impl BufRead,
max_bytes: usize,
) -> McpResult<Option<SseEvent>> {
let mut parser = SseParser::default();
let mut line = String::new();
let mut total = 0usize;
loop {
line.clear();
let read = reader.read_line(&mut line).map_err(McpError::transport)?;
if read == 0 {
return Ok(parser.dispatch());
}
total = total.saturating_add(read);
if total > max_bytes {
return Err(McpError::Transport(format!(
"MCP SSE event exceeded {max_bytes} bytes"
)));
}
while matches!(line.chars().last(), Some('\n' | '\r')) {
line.pop();
}
if let Some(event) = parser.push_line(&line) {
return Ok(Some(event));
}
}
}
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 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}");
}
}