use bytes::Bytes;
use futures::{Stream, StreamExt};
use tracing::debug;
use crate::error::{LambdaError, Result};
pub fn format_sse_event(data: &str, event_type: Option<&str>, event_id: Option<&str>) -> String {
let mut event = String::new();
if let Some(id) = event_id {
event.push_str(&format!("id: {}\n", id));
}
if let Some(event_type) = event_type {
event.push_str(&format!("event: {}\n", event_type));
}
for line in data.lines() {
event.push_str(&format!("data: {}\n", line));
}
event.push('\n'); event
}
pub fn create_sse_stream<T>(
events: Vec<T>,
formatter: impl Fn(&T) -> String + Send + 'static,
) -> impl Stream<Item = Result<Bytes>> + Send + 'static
where
T: Send + 'static,
{
async_stream::stream! {
for event in events {
let sse_data = formatter(&event);
let bytes = Bytes::from(sse_data);
yield Ok(bytes);
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
}
}
}
pub fn create_heartbeat_stream(
interval_secs: u64,
) -> impl Stream<Item = Result<Bytes>> + Send + 'static {
async_stream::stream! {
let mut interval = tokio::time::interval(
tokio::time::Duration::from_secs(interval_secs)
);
loop {
interval.tick().await;
let heartbeat = format_sse_event(
"heartbeat",
Some("heartbeat"),
Some(&chrono::Utc::now().timestamp().to_string())
);
yield Ok(Bytes::from(heartbeat));
}
}
}
pub fn merge_sse_streams<S1, S2>(
stream1: S1,
stream2: S2,
) -> impl Stream<Item = Result<Bytes>> + Send + 'static
where
S1: Stream<Item = Result<Bytes>> + Send + 'static,
S2: Stream<Item = Result<Bytes>> + Send + 'static,
{
use futures::stream::select;
select(stream1.map(|item| (1, item)), stream2.map(|item| (2, item))).map(|(_, result)| result)
}
pub fn validate_sse_event(event: &str) -> Result<()> {
if event.contains('\0') {
return Err(LambdaError::Sse(
"SSE events cannot contain null bytes".to_string(),
));
}
if event.len() > 1_048_576 {
debug!("Warning: SSE event is very large ({} bytes)", event.len());
}
if event.contains('\r') && !event.contains("\r\n") {
return Err(LambdaError::Sse(
"SSE events should use LF or CRLF line endings, not standalone CR".to_string(),
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use futures::stream;
#[test]
fn test_format_sse_event() {
let event = format_sse_event("Hello, World!", Some("message"), Some("123"));
assert!(event.contains("id: 123\n"));
assert!(event.contains("event: message\n"));
assert!(event.contains("data: Hello, World!\n"));
assert!(event.ends_with("\n\n"));
}
#[test]
fn test_format_multiline_event() {
let data = "Line 1\nLine 2\nLine 3";
let event = format_sse_event(data, None, None);
assert!(event.contains("data: Line 1\n"));
assert!(event.contains("data: Line 2\n"));
assert!(event.contains("data: Line 3\n"));
}
#[tokio::test]
async fn test_create_sse_stream() {
use futures::StreamExt;
use futures::pin_mut;
let events = vec!["event1", "event2", "event3"];
let stream = create_sse_stream(events, |s| format_sse_event(s, Some("test"), None));
pin_mut!(stream);
let first_event = stream.next().await.unwrap().unwrap();
let event_str = String::from_utf8(first_event.to_vec()).unwrap();
assert!(event_str.contains("event: test\n"));
assert!(event_str.contains("data: event1\n"));
}
#[test]
fn test_validate_sse_event() {
assert!(validate_sse_event("Normal event").is_ok());
assert!(validate_sse_event("Event\nwith\nnewlines").is_ok());
assert!(validate_sse_event("Event with\0null byte").is_err());
assert!(validate_sse_event("Event with\rstandalone CR").is_err());
assert!(validate_sse_event("Event with\r\nCRLF").is_ok());
}
#[tokio::test]
async fn test_merge_streams() {
let stream1 = stream::iter(vec![
Ok(Bytes::from("stream1-1")),
Ok(Bytes::from("stream1-2")),
]);
let stream2 = stream::iter(vec![
Ok(Bytes::from("stream2-1")),
Ok(Bytes::from("stream2-2")),
]);
let merged = merge_sse_streams(stream1, stream2);
let results: Vec<_> = merged.collect().await;
assert_eq!(results.len(), 4);
}
}