Skip to main content

dcp/compat/
sse.rs

1//! Server-Sent Events (SSE) transport for MCP compatibility.
2//!
3//! Provides HTTP SSE handling and event stream formatting for MCP clients.
4//! Supports reconnection with Last-Event-ID for event replay.
5
6use std::collections::VecDeque;
7use std::fmt::Write;
8use std::sync::atomic::{AtomicU64, Ordering};
9use std::sync::{Arc, RwLock};
10
11use crate::security::{sanitize_json_value, sanitize_text};
12
13use super::json_rpc::{JsonRpcParser, JsonRpcRequest, RequestId, DEFAULT_MAX_JSONRPC_REQUEST_SIZE};
14
15/// Default maximum number of queued SSE POST messages.
16pub const DEFAULT_MAX_SSE_PENDING_RESPONSES: usize = 1024;
17/// Default maximum aggregate bytes retained in the SSE POST queue.
18pub const DEFAULT_MAX_SSE_PENDING_RESPONSE_BYTES: usize =
19    DEFAULT_MAX_JSONRPC_REQUEST_SIZE * DEFAULT_MAX_SSE_PENDING_RESPONSES;
20
21/// Errors returned by checked SSE replay.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum SseReplayError {
24    /// Requested event ID has already been evicted from the replay buffer.
25    StaleEventId,
26    /// Requested event ID is newer than the latest retained event.
27    FutureEventId,
28}
29
30/// SSE event types
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub enum SseEventType {
33    /// JSON-RPC message event
34    Message,
35    /// Endpoint information event
36    Endpoint,
37    /// Error event
38    Error,
39    /// Ping/keepalive event
40    Ping,
41}
42
43impl SseEventType {
44    /// Get the event type string
45    pub fn as_str(&self) -> &'static str {
46        match self {
47            Self::Message => "message",
48            Self::Endpoint => "endpoint",
49            Self::Error => "error",
50            Self::Ping => "ping",
51        }
52    }
53}
54
55/// SSE event for streaming
56#[derive(Debug, Clone)]
57pub struct SseEvent {
58    /// Event type
59    pub event_type: SseEventType,
60    /// Event data (JSON payload)
61    pub data: String,
62    /// Optional event ID
63    pub id: Option<String>,
64    /// Optional retry interval in milliseconds
65    pub retry: Option<u32>,
66}
67
68impl SseEvent {
69    /// Create a new message event
70    pub fn message(data: String) -> Self {
71        Self {
72            event_type: SseEventType::Message,
73            data,
74            id: None,
75            retry: None,
76        }
77    }
78
79    /// Create a new endpoint event
80    pub fn endpoint(url: String) -> Self {
81        Self {
82            event_type: SseEventType::Endpoint,
83            data: url,
84            id: None,
85            retry: None,
86        }
87    }
88
89    /// Create a new error event
90    pub fn error(message: String) -> Self {
91        Self {
92            event_type: SseEventType::Error,
93            data: sanitize_text(&message),
94            id: None,
95            retry: None,
96        }
97    }
98
99    /// Create a ping event
100    pub fn ping() -> Self {
101        Self {
102            event_type: SseEventType::Ping,
103            data: String::new(),
104            id: None,
105            retry: None,
106        }
107    }
108
109    /// Set the event ID
110    pub fn with_id(mut self, id: String) -> Self {
111        self.id = Some(sanitize_sse_field(&id));
112        self
113    }
114
115    /// Set the retry interval
116    pub fn with_retry(mut self, retry_ms: u32) -> Self {
117        self.retry = Some(retry_ms);
118        self
119    }
120
121    /// Format the event as SSE text
122    pub fn format(&self) -> String {
123        let mut output = String::new();
124
125        // Event type
126        let _ = writeln!(output, "event: {}", self.event_type.as_str());
127
128        // Event ID if present
129        if let Some(ref id) = self.id {
130            let _ = writeln!(output, "id: {}", sanitize_sse_field(id));
131        }
132
133        // Retry if present
134        if let Some(retry) = self.retry {
135            let _ = writeln!(output, "retry: {}", retry);
136        }
137
138        // Data lines (split by newlines)
139        let data = sanitize_sse_data(&self.data);
140        for line in data.lines() {
141            let _ = writeln!(output, "data: {}", line);
142        }
143
144        // Empty data line if no data
145        if data.is_empty() {
146            let _ = writeln!(output, "data:");
147        }
148
149        // Blank line to end event
150        let _ = writeln!(output);
151
152        output
153    }
154
155    /// Get the numeric ID if present
156    pub fn numeric_id(&self) -> Option<u64> {
157        self.id.as_ref().and_then(|id| id.parse().ok())
158    }
159}
160
161/// Buffered event for replay
162#[derive(Debug, Clone)]
163struct BufferedEvent {
164    /// Event ID
165    id: u64,
166    /// Formatted event data
167    formatted: String,
168}
169
170fn sanitize_replay_payload(json: &str) -> String {
171    serde_json::from_str::<serde_json::Value>(json)
172        .ok()
173        .and_then(|value| {
174            let sanitized = sanitize_json_value(&value);
175            if sanitized == value {
176                Some(json.to_string())
177            } else {
178                serde_json::to_string(&sanitized).ok()
179            }
180        })
181        .unwrap_or_else(|| sanitize_text(json))
182}
183
184fn sanitize_sse_field(value: &str) -> String {
185    sanitize_text(value)
186        .chars()
187        .map(|ch| if matches!(ch, '\r' | '\n') { ' ' } else { ch })
188        .collect()
189}
190
191fn sanitize_sse_data(value: &str) -> String {
192    sanitize_replay_payload(value)
193        .replace("\r\n", "\n")
194        .replace('\r', "\n")
195}
196
197/// Event buffer for reconnection replay
198#[derive(Debug)]
199pub struct EventBuffer {
200    /// Buffered events
201    events: VecDeque<BufferedEvent>,
202    /// Maximum buffer size
203    max_size: usize,
204}
205
206impl EventBuffer {
207    /// Create a new event buffer
208    pub fn new(max_size: usize) -> Self {
209        Self {
210            events: VecDeque::with_capacity(max_size.min(1000)),
211            max_size,
212        }
213    }
214
215    /// Add an event to the buffer
216    pub fn push(&mut self, id: u64, formatted: String) {
217        if self.max_size == 0 {
218            return;
219        }
220
221        // Remove oldest if at capacity
222        while self.events.len() >= self.max_size {
223            self.events.pop_front();
224        }
225        self.events.push_back(BufferedEvent { id, formatted });
226    }
227
228    /// Get events after a given ID for replay
229    pub fn events_after(&self, last_id: u64) -> Vec<String> {
230        self.events
231            .iter()
232            .filter(|e| e.id > last_id)
233            .map(|e| e.formatted.clone())
234            .collect()
235    }
236
237    /// Get the latest event ID
238    pub fn latest_id(&self) -> Option<u64> {
239        self.events.back().map(|e| e.id)
240    }
241
242    /// Get the oldest retained event ID.
243    pub fn oldest_id(&self) -> Option<u64> {
244        self.events.front().map(|e| e.id)
245    }
246
247    /// Get buffer size
248    pub fn len(&self) -> usize {
249        self.events.len()
250    }
251
252    /// Check if buffer is empty
253    pub fn is_empty(&self) -> bool {
254        self.events.is_empty()
255    }
256
257    /// Clear the buffer
258    pub fn clear(&mut self) {
259        self.events.clear();
260    }
261}
262
263/// SSE transport handler with reconnection support
264pub struct SseTransport {
265    /// Event counter for IDs
266    event_counter: AtomicU64,
267    /// Default retry interval
268    default_retry_ms: u32,
269    /// Event buffer for replay
270    buffer: Arc<RwLock<EventBuffer>>,
271}
272
273impl Default for SseTransport {
274    fn default() -> Self {
275        Self::new()
276    }
277}
278
279impl SseTransport {
280    /// Create a new SSE transport
281    pub fn new() -> Self {
282        Self::with_buffer_size(100)
283    }
284
285    /// Create with custom buffer size
286    pub fn with_buffer_size(buffer_size: usize) -> Self {
287        Self {
288            event_counter: AtomicU64::new(0),
289            default_retry_ms: 3000,
290            buffer: Arc::new(RwLock::new(EventBuffer::new(buffer_size))),
291        }
292    }
293
294    /// Set the default retry interval
295    pub fn with_retry(mut self, retry_ms: u32) -> Self {
296        self.default_retry_ms = retry_ms;
297        self
298    }
299
300    /// Get the content type header for SSE
301    pub fn content_type() -> &'static str {
302        "text/event-stream"
303    }
304
305    /// Get required headers for SSE response
306    pub fn headers() -> Vec<(&'static str, &'static str)> {
307        vec![
308            ("Content-Type", "text/event-stream"),
309            ("Cache-Control", "no-cache"),
310            ("Connection", "keep-alive"),
311        ]
312    }
313
314    /// Get the next event ID
315    fn next_id(&self) -> u64 {
316        self.event_counter.fetch_add(1, Ordering::SeqCst) + 1
317    }
318
319    /// Format a JSON-RPC response as an SSE event
320    pub fn format_response(&self, json: &str) -> String {
321        let id = self.next_id();
322        let event = SseEvent::message(json.to_string())
323            .with_id(id.to_string())
324            .with_retry(self.default_retry_ms);
325        let formatted = event.format();
326
327        let replay_event = SseEvent::message(sanitize_replay_payload(json))
328            .with_id(id.to_string())
329            .with_retry(self.default_retry_ms);
330        let replay_formatted = replay_event.format();
331
332        // Buffer a sanitized copy for replay.
333        if let Ok(mut buffer) = self.buffer.write() {
334            buffer.push(id, replay_formatted);
335        }
336
337        formatted
338    }
339
340    /// Format an error as an SSE event
341    pub fn format_error(&self, message: &str) -> String {
342        let id = self.next_id();
343        let event = SseEvent::error(sanitize_text(message)).with_id(id.to_string());
344        let formatted = event.format();
345
346        // Buffer for replay
347        if let Ok(mut buffer) = self.buffer.write() {
348            buffer.push(id, formatted.clone());
349        }
350
351        formatted
352    }
353
354    /// Format a ping/keepalive event (not buffered)
355    pub fn format_ping(&self) -> String {
356        SseEvent::ping().format()
357    }
358
359    /// Format an endpoint event (for initial connection)
360    pub fn format_endpoint(&self, url: &str) -> String {
361        let id = self.next_id();
362        let event = SseEvent::endpoint(sanitize_text(url)).with_id(id.to_string());
363        event.format()
364    }
365
366    /// Parse Last-Event-ID header to resume from
367    pub fn parse_last_event_id(header: &str) -> Option<u64> {
368        header.trim().parse().ok()
369    }
370
371    /// Get events for replay after reconnection
372    pub fn get_replay_events(&self, last_event_id: u64) -> Vec<String> {
373        if let Ok(buffer) = self.buffer.read() {
374            buffer.events_after(last_event_id)
375        } else {
376            Vec::new()
377        }
378    }
379
380    /// Get replay events with explicit stale/future cursor handling.
381    pub fn checked_replay_events(&self, last_event_id: u64) -> Result<Vec<String>, SseReplayError> {
382        let buffer = self
383            .buffer
384            .read()
385            .map_err(|_| SseReplayError::StaleEventId)?;
386
387        if buffer.is_empty() {
388            let current = self.current_event_id();
389            if last_event_id > current {
390                return Err(SseReplayError::FutureEventId);
391            }
392            if current > 0 && last_event_id < current {
393                return Err(SseReplayError::StaleEventId);
394            }
395            return Ok(Vec::new());
396        }
397
398        if let Some(latest) = buffer.latest_id() {
399            if last_event_id > latest {
400                return Err(SseReplayError::FutureEventId);
401            }
402        }
403
404        if let Some(oldest) = buffer.oldest_id() {
405            if last_event_id < oldest.saturating_sub(1) {
406                return Err(SseReplayError::StaleEventId);
407            }
408        }
409
410        Ok(buffer.events_after(last_event_id))
411    }
412
413    /// Get the current event counter value
414    pub fn current_event_id(&self) -> u64 {
415        self.event_counter.load(Ordering::SeqCst)
416    }
417
418    /// Get buffer statistics
419    pub fn buffer_stats(&self) -> (usize, Option<u64>) {
420        if let Ok(buffer) = self.buffer.read() {
421            (buffer.len(), buffer.latest_id())
422        } else {
423            (0, None)
424        }
425    }
426
427    /// Reset event counter and clear buffer (for testing)
428    pub fn reset(&self) {
429        self.event_counter.store(0, Ordering::SeqCst);
430        if let Ok(mut buffer) = self.buffer.write() {
431            buffer.clear();
432        }
433    }
434}
435
436/// HTTP POST message handler for SSE transport
437pub struct SseMessageHandler {
438    /// Pending responses to send
439    responses: Arc<RwLock<VecDeque<String>>>,
440    /// Maximum accepted JSON-RPC request size in bytes.
441    max_message_size: usize,
442    /// Maximum number of queued pending responses.
443    max_pending_responses: usize,
444    /// Maximum aggregate bytes queued in pending responses.
445    max_pending_response_bytes: usize,
446}
447
448impl Default for SseMessageHandler {
449    fn default() -> Self {
450        Self::new()
451    }
452}
453
454impl SseMessageHandler {
455    /// Create a new message handler
456    pub fn new() -> Self {
457        Self {
458            responses: Arc::new(RwLock::new(VecDeque::new())),
459            max_message_size: DEFAULT_MAX_JSONRPC_REQUEST_SIZE,
460            max_pending_responses: DEFAULT_MAX_SSE_PENDING_RESPONSES,
461            max_pending_response_bytes: DEFAULT_MAX_SSE_PENDING_RESPONSE_BYTES,
462        }
463    }
464
465    /// Set maximum accepted JSON-RPC request size for POST messages.
466    pub fn with_max_message_size(mut self, max_message_size: usize) -> Self {
467        self.max_message_size = max_message_size;
468        self
469    }
470
471    /// Set maximum number of queued POST messages.
472    pub fn with_max_pending_responses(mut self, max_pending_responses: usize) -> Self {
473        self.max_pending_responses = max_pending_responses;
474        self
475    }
476
477    /// Set maximum aggregate bytes queued by POST messages.
478    pub fn with_max_pending_response_bytes(mut self, max_pending_response_bytes: usize) -> Self {
479        self.max_pending_response_bytes = max_pending_response_bytes;
480        self
481    }
482
483    /// Handle an incoming POST message
484    pub fn handle_message(&self, json: &str) -> Result<(), String> {
485        let request = JsonRpcParser::parse_request_with_limit(json, self.max_message_size)
486            .map_err(|_| "Invalid JSON-RPC request".to_string())?;
487        if !is_supported_sse_post_method(&request.method) {
488            return Err("Unsupported JSON-RPC method on SSE POST".to_string());
489        }
490        validate_supported_sse_notification(&request)?;
491        let notification_only = is_supported_sse_notification(&request.method);
492        if request.is_notification() && !notification_only {
493            return Err("JSON-RPC notification not accepted on SSE POST".to_string());
494        }
495        if !request.is_notification() && notification_only {
496            return Err("JSON-RPC notification method must not include id".to_string());
497        }
498        if request.is_notification() {
499            return Ok(());
500        }
501
502        // Queue response (in real implementation, this would process the request).
503        let mut responses = self
504            .responses
505            .write()
506            .map_err(|_| "SSE response queue unavailable".to_string())?;
507        if responses.len() >= self.max_pending_responses {
508            return Err("SSE response queue full".to_string());
509        }
510        let sanitized = sanitize_replay_payload(json);
511        let queued_bytes = responses.iter().fold(0usize, |total, response| {
512            total.saturating_add(response.len())
513        });
514        if queued_bytes.saturating_add(sanitized.len()) > self.max_pending_response_bytes {
515            return Err("SSE response queue byte budget exceeded".to_string());
516        }
517        responses.push_back(sanitized);
518
519        Ok(())
520    }
521
522    /// Get pending responses
523    pub fn take_responses(&self) -> Vec<String> {
524        if let Ok(mut responses) = self.responses.write() {
525            responses.drain(..).collect()
526        } else {
527            Vec::new()
528        }
529    }
530
531    /// Check if there are pending responses
532    pub fn has_pending(&self) -> bool {
533        if let Ok(responses) = self.responses.read() {
534            !responses.is_empty()
535        } else {
536            false
537        }
538    }
539}
540
541fn is_supported_sse_notification(method: &str) -> bool {
542    matches!(
543        method,
544        "notifications/initialized" | "notifications/cancelled"
545    )
546}
547
548fn validate_supported_sse_notification(request: &JsonRpcRequest) -> Result<(), String> {
549    if request.method == "notifications/initialized" {
550        return if request.params.is_none() {
551            Ok(())
552        } else {
553            Err("Invalid JSON-RPC notification params".to_string())
554        };
555    }
556
557    if request.method != "notifications/cancelled" {
558        return Ok(());
559    }
560
561    let params = request
562        .params
563        .as_ref()
564        .and_then(|params| params.as_object())
565        .ok_or_else(|| "Invalid JSON-RPC notification params".to_string())?;
566    if params
567        .keys()
568        .any(|key| key != "requestId" && key != "reason")
569    {
570        return Err("Invalid JSON-RPC notification params".to_string());
571    }
572
573    let request_id = params
574        .get("requestId")
575        .ok_or_else(|| "Invalid JSON-RPC notification params".to_string())?;
576
577    RequestId::try_from_json_value(request_id)
578        .map_err(|_| "Invalid JSON-RPC notification params".to_string())?;
579
580    if params
581        .get("reason")
582        .is_some_and(|reason| !reason.is_string())
583    {
584        return Err("Invalid JSON-RPC notification params".to_string());
585    }
586
587    Ok(())
588}
589
590fn is_supported_sse_post_method(method: &str) -> bool {
591    matches!(
592        method,
593        "initialize"
594            | "ping"
595            | "tools/list"
596            | "tools/call"
597            | "resources/list"
598            | "resources/read"
599            | "resources/subscribe"
600            | "resources/unsubscribe"
601            | "prompts/list"
602            | "prompts/get"
603            | "logging/setLevel"
604            | "completion/complete"
605    ) || is_supported_sse_notification(method)
606}
607
608/// SSE endpoint configuration
609#[derive(Debug, Clone)]
610pub struct SseEndpointConfig {
611    /// Path for SSE event stream (GET)
612    pub events_path: String,
613    /// Path for client messages (POST)
614    pub messages_path: String,
615    /// Ping interval in seconds
616    pub ping_interval_secs: u64,
617    /// Event buffer size
618    pub buffer_size: usize,
619    /// Default retry interval in milliseconds
620    pub retry_ms: u32,
621}
622
623impl Default for SseEndpointConfig {
624    fn default() -> Self {
625        Self {
626            events_path: "/events".to_string(),
627            messages_path: "/message".to_string(),
628            ping_interval_secs: 30,
629            buffer_size: 100,
630            retry_ms: 3000,
631        }
632    }
633}
634
635#[cfg(test)]
636mod tests {
637    use super::*;
638
639    #[test]
640    fn test_sse_event_format_message() {
641        let event = SseEvent::message(r#"{"jsonrpc":"2.0","result":{},"id":1}"#.to_string());
642        let formatted = event.format();
643
644        assert!(formatted.contains("event: message"));
645        assert!(formatted.contains(r#"data: {"jsonrpc":"2.0","result":{},"id":1}"#));
646        assert!(formatted.ends_with("\n\n"));
647    }
648
649    #[test]
650    fn test_sse_event_with_id_and_retry() {
651        let event = SseEvent::message("test".to_string())
652            .with_id("42".to_string())
653            .with_retry(5000);
654        let formatted = event.format();
655
656        assert!(formatted.contains("id: 42"));
657        assert!(formatted.contains("retry: 5000"));
658    }
659
660    #[test]
661    fn test_sse_event_multiline_data() {
662        let event = SseEvent::message("line1\nline2\nline3".to_string());
663        let formatted = event.format();
664
665        assert!(formatted.contains("data: line1"));
666        assert!(formatted.contains("data: line2"));
667        assert!(formatted.contains("data: line3"));
668    }
669
670    #[test]
671    fn test_sse_transport_format_response() {
672        let transport = SseTransport::new();
673        let response = transport.format_response(r#"{"result":"ok"}"#);
674
675        assert!(response.contains("event: message"));
676        assert!(response.contains("id: 1"));
677        assert!(response.contains("retry: 3000"));
678        assert!(response.contains(r#"data: {"result":"ok"}"#));
679    }
680
681    #[test]
682    fn test_sse_transport_format_error() {
683        let transport = SseTransport::new();
684        let error = transport.format_error("Something went wrong");
685
686        assert!(error.contains("event: error"));
687        assert!(error.contains("data: Something went wrong"));
688    }
689
690    #[test]
691    fn test_sse_transport_format_ping() {
692        let transport = SseTransport::new();
693        let ping = transport.format_ping();
694
695        assert!(ping.contains("event: ping"));
696        assert!(ping.contains("data:"));
697    }
698
699    #[test]
700    fn test_sse_transport_headers() {
701        let headers = SseTransport::headers();
702
703        assert!(headers.contains(&("Content-Type", "text/event-stream")));
704        assert!(headers.contains(&("Cache-Control", "no-cache")));
705        assert!(headers.contains(&("Connection", "keep-alive")));
706    }
707
708    #[test]
709    fn test_parse_last_event_id() {
710        assert_eq!(SseTransport::parse_last_event_id("42"), Some(42));
711        assert_eq!(SseTransport::parse_last_event_id("  100  "), Some(100));
712        assert_eq!(SseTransport::parse_last_event_id("invalid"), None);
713    }
714
715    #[test]
716    fn test_sse_event_types() {
717        assert_eq!(SseEventType::Message.as_str(), "message");
718        assert_eq!(SseEventType::Endpoint.as_str(), "endpoint");
719        assert_eq!(SseEventType::Error.as_str(), "error");
720        assert_eq!(SseEventType::Ping.as_str(), "ping");
721    }
722
723    #[test]
724    fn test_event_buffer() {
725        let mut buffer = EventBuffer::new(3);
726
727        buffer.push(1, "event1".to_string());
728        buffer.push(2, "event2".to_string());
729        buffer.push(3, "event3".to_string());
730
731        assert_eq!(buffer.len(), 3);
732        assert_eq!(buffer.latest_id(), Some(3));
733
734        // Get events after ID 1
735        let events = buffer.events_after(1);
736        assert_eq!(events.len(), 2);
737        assert_eq!(events[0], "event2");
738        assert_eq!(events[1], "event3");
739
740        // Add one more, should evict oldest
741        buffer.push(4, "event4".to_string());
742        assert_eq!(buffer.len(), 3);
743        assert_eq!(buffer.latest_id(), Some(4));
744
745        // Event 1 should be gone
746        let events = buffer.events_after(0);
747        assert_eq!(events.len(), 3);
748        assert!(!events.contains(&"event1".to_string()));
749    }
750
751    #[test]
752    fn test_sse_transport_replay() {
753        let transport = SseTransport::with_buffer_size(10);
754
755        // Send some events
756        transport.format_response(r#"{"id":1}"#);
757        transport.format_response(r#"{"id":2}"#);
758        transport.format_response(r#"{"id":3}"#);
759
760        // Get replay events after ID 1
761        let replay = transport.get_replay_events(1);
762        assert_eq!(replay.len(), 2);
763
764        // Get all events
765        let all = transport.get_replay_events(0);
766        assert_eq!(all.len(), 3);
767    }
768
769    #[test]
770    fn test_sse_message_handler() {
771        let handler = SseMessageHandler::new();
772
773        // Handle valid message
774        assert!(handler
775            .handle_message(r#"{"jsonrpc":"2.0","method":"ping","id":1}"#)
776            .is_ok());
777        assert!(handler.has_pending());
778
779        // Handle invalid message
780        assert!(handler.handle_message("not json").is_err());
781
782        // Take responses
783        let responses = handler.take_responses();
784        assert_eq!(responses.len(), 1);
785        assert!(!handler.has_pending());
786    }
787}