1use 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
15pub const DEFAULT_MAX_SSE_PENDING_RESPONSES: usize = 1024;
17pub const DEFAULT_MAX_SSE_PENDING_RESPONSE_BYTES: usize =
19 DEFAULT_MAX_JSONRPC_REQUEST_SIZE * DEFAULT_MAX_SSE_PENDING_RESPONSES;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum SseReplayError {
24 StaleEventId,
26 FutureEventId,
28}
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub enum SseEventType {
33 Message,
35 Endpoint,
37 Error,
39 Ping,
41}
42
43impl SseEventType {
44 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#[derive(Debug, Clone)]
57pub struct SseEvent {
58 pub event_type: SseEventType,
60 pub data: String,
62 pub id: Option<String>,
64 pub retry: Option<u32>,
66}
67
68impl SseEvent {
69 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 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 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 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 pub fn with_id(mut self, id: String) -> Self {
111 self.id = Some(sanitize_sse_field(&id));
112 self
113 }
114
115 pub fn with_retry(mut self, retry_ms: u32) -> Self {
117 self.retry = Some(retry_ms);
118 self
119 }
120
121 pub fn format(&self) -> String {
123 let mut output = String::new();
124
125 let _ = writeln!(output, "event: {}", self.event_type.as_str());
127
128 if let Some(ref id) = self.id {
130 let _ = writeln!(output, "id: {}", sanitize_sse_field(id));
131 }
132
133 if let Some(retry) = self.retry {
135 let _ = writeln!(output, "retry: {}", retry);
136 }
137
138 let data = sanitize_sse_data(&self.data);
140 for line in data.lines() {
141 let _ = writeln!(output, "data: {}", line);
142 }
143
144 if data.is_empty() {
146 let _ = writeln!(output, "data:");
147 }
148
149 let _ = writeln!(output);
151
152 output
153 }
154
155 pub fn numeric_id(&self) -> Option<u64> {
157 self.id.as_ref().and_then(|id| id.parse().ok())
158 }
159}
160
161#[derive(Debug, Clone)]
163struct BufferedEvent {
164 id: u64,
166 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#[derive(Debug)]
199pub struct EventBuffer {
200 events: VecDeque<BufferedEvent>,
202 max_size: usize,
204}
205
206impl EventBuffer {
207 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 pub fn push(&mut self, id: u64, formatted: String) {
217 if self.max_size == 0 {
218 return;
219 }
220
221 while self.events.len() >= self.max_size {
223 self.events.pop_front();
224 }
225 self.events.push_back(BufferedEvent { id, formatted });
226 }
227
228 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 pub fn latest_id(&self) -> Option<u64> {
239 self.events.back().map(|e| e.id)
240 }
241
242 pub fn oldest_id(&self) -> Option<u64> {
244 self.events.front().map(|e| e.id)
245 }
246
247 pub fn len(&self) -> usize {
249 self.events.len()
250 }
251
252 pub fn is_empty(&self) -> bool {
254 self.events.is_empty()
255 }
256
257 pub fn clear(&mut self) {
259 self.events.clear();
260 }
261}
262
263pub struct SseTransport {
265 event_counter: AtomicU64,
267 default_retry_ms: u32,
269 buffer: Arc<RwLock<EventBuffer>>,
271}
272
273impl Default for SseTransport {
274 fn default() -> Self {
275 Self::new()
276 }
277}
278
279impl SseTransport {
280 pub fn new() -> Self {
282 Self::with_buffer_size(100)
283 }
284
285 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 pub fn with_retry(mut self, retry_ms: u32) -> Self {
296 self.default_retry_ms = retry_ms;
297 self
298 }
299
300 pub fn content_type() -> &'static str {
302 "text/event-stream"
303 }
304
305 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 fn next_id(&self) -> u64 {
316 self.event_counter.fetch_add(1, Ordering::SeqCst) + 1
317 }
318
319 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 if let Ok(mut buffer) = self.buffer.write() {
334 buffer.push(id, replay_formatted);
335 }
336
337 formatted
338 }
339
340 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 if let Ok(mut buffer) = self.buffer.write() {
348 buffer.push(id, formatted.clone());
349 }
350
351 formatted
352 }
353
354 pub fn format_ping(&self) -> String {
356 SseEvent::ping().format()
357 }
358
359 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 pub fn parse_last_event_id(header: &str) -> Option<u64> {
368 header.trim().parse().ok()
369 }
370
371 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 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 pub fn current_event_id(&self) -> u64 {
415 self.event_counter.load(Ordering::SeqCst)
416 }
417
418 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 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
436pub struct SseMessageHandler {
438 responses: Arc<RwLock<VecDeque<String>>>,
440 max_message_size: usize,
442 max_pending_responses: usize,
444 max_pending_response_bytes: usize,
446}
447
448impl Default for SseMessageHandler {
449 fn default() -> Self {
450 Self::new()
451 }
452}
453
454impl SseMessageHandler {
455 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 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 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 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 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 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 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 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#[derive(Debug, Clone)]
610pub struct SseEndpointConfig {
611 pub events_path: String,
613 pub messages_path: String,
615 pub ping_interval_secs: u64,
617 pub buffer_size: usize,
619 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 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 buffer.push(4, "event4".to_string());
742 assert_eq!(buffer.len(), 3);
743 assert_eq!(buffer.latest_id(), Some(4));
744
745 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 transport.format_response(r#"{"id":1}"#);
757 transport.format_response(r#"{"id":2}"#);
758 transport.format_response(r#"{"id":3}"#);
759
760 let replay = transport.get_replay_events(1);
762 assert_eq!(replay.len(), 2);
763
764 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 assert!(handler
775 .handle_message(r#"{"jsonrpc":"2.0","method":"ping","id":1}"#)
776 .is_ok());
777 assert!(handler.has_pending());
778
779 assert!(handler.handle_message("not json").is_err());
781
782 let responses = handler.take_responses();
784 assert_eq!(responses.len(), 1);
785 assert!(!handler.has_pending());
786 }
787}