use std::pin::Pin;
use std::time::Duration;
use bytes::Bytes;
use futures::StreamExt;
use futures::stream::Stream;
use super::events::SessionEvent;
use crate::{Error, Result};
const DEFAULT_SSE_TIMEOUT_SECS: u64 = 300;
pub fn default_sse_timeout() -> Duration {
Duration::from_secs(DEFAULT_SSE_TIMEOUT_SECS)
}
pub fn process_managed_agents_sse<S>(
byte_stream: S,
timeout: Duration,
) -> Pin<Box<dyn Stream<Item = Result<SessionEvent>> + Send>>
where
S: Stream<Item = std::result::Result<Bytes, reqwest::Error>> + Unpin + Send + 'static,
{
let state = SseState {
byte_stream: Box::pin(byte_stream),
buffer: String::new(),
timeout,
done: false,
};
Box::pin(futures::stream::unfold(state, |mut state| async move {
if state.done {
return None;
}
loop {
if let Some(event_result) = state.try_parse_next_event() {
return Some((event_result, state));
}
let read_result = tokio::time::timeout(state.timeout, state.byte_stream.next()).await;
match read_result {
Ok(Some(Ok(bytes))) => {
match std::str::from_utf8(&bytes) {
Ok(text) => {
state.buffer.push_str(text);
}
Err(e) => {
let err = Error::Encoding {
message: format!("invalid UTF-8 in SSE stream: {e}"),
source: None,
};
return Some((Err(err), state));
}
}
}
Ok(Some(Err(e))) => {
state.done = true;
let err = Error::Connection {
message: format!("SSE stream error: {e}"),
source: None,
};
return Some((Err(err), state));
}
Ok(None) => {
if let Some(event_result) = state.try_parse_next_event() {
state.done = true;
return Some((event_result, state));
}
return None;
}
Err(_elapsed) => {
let err = Error::Timeout {
message: format!(
"no data received on SSE stream within {} seconds",
state.timeout.as_secs()
),
duration: Some(state.timeout.as_secs_f64()),
};
return Some((Err(err), state));
}
}
}
}))
}
struct SseState {
byte_stream: Pin<Box<dyn Stream<Item = std::result::Result<Bytes, reqwest::Error>> + Send>>,
buffer: String,
timeout: Duration,
done: bool,
}
impl SseState {
fn try_parse_next_event(&mut self) -> Option<Result<SessionEvent>> {
loop {
let delimiter_pos = self.buffer.find("\n\n")?;
let event_block = self.buffer[..delimiter_pos].to_string();
self.buffer = self.buffer[delimiter_pos + 2..].to_string();
if event_block.trim().is_empty() {
continue;
}
match parse_sse_event(&event_block) {
Some(Ok(session_event)) => return Some(Ok(session_event)),
Some(Err(err)) => {
tracing::warn!(
error = %err,
event_block = %event_block,
"skipping SSE event with invalid JSON"
);
continue;
}
None => {
continue;
}
}
}
}
}
fn parse_sse_event(block: &str) -> Option<Result<SessionEvent>> {
let mut data_lines: Vec<&str> = Vec::new();
for line in block.lines() {
if let Some(data) = line.strip_prefix("data:") {
data_lines.push(data.trim_start());
}
}
if data_lines.is_empty() {
return None;
}
let data = data_lines.join("\n");
if data.is_empty() {
return None;
}
match serde_json::from_str::<SessionEvent>(&data) {
Ok(event) => Some(Ok(event)),
Err(e) => Some(Err(Error::Serialization {
message: format!("failed to deserialize SSE event data: {e}"),
source: None,
})),
}
}
#[cfg(test)]
mod tests {
use super::*;
use bytes::Bytes;
use futures::stream;
fn byte_stream_from_chunks(
chunks: Vec<&str>,
) -> impl Stream<Item = std::result::Result<Bytes, reqwest::Error>> + Unpin {
stream::iter(chunks.into_iter().map(|s| Ok(Bytes::from(s.to_string()))).collect::<Vec<_>>())
}
#[tokio::test]
async fn test_parse_single_agent_message_event() {
let chunks = vec![
"event: agent.message\ndata: {\"type\":\"agent.message\",\"content\":\"Hello!\"}\n\n",
];
let stream = byte_stream_from_chunks(chunks);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentMessage { id: None, content: serde_json::json!("Hello!") }
);
assert!(event_stream.next().await.is_none());
}
#[tokio::test]
async fn test_parse_multiple_events() {
let chunks = vec![
"event: session.status_running\ndata: {\"type\":\"session.status_running\"}\n\n",
"event: agent.message\ndata: {\"type\":\"agent.message\",\"content\":\"Hi\"}\n\n",
"event: session.status_idle\ndata: {\"type\":\"session.status_idle\"}\n\n",
];
let stream = byte_stream_from_chunks(chunks);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event1 = event_stream.next().await.unwrap().unwrap();
assert_eq!(event1, SessionEvent::StatusRunning {});
let event2 = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event2,
SessionEvent::AgentMessage { id: None, content: serde_json::json!("Hi") }
);
let event3 = event_stream.next().await.unwrap().unwrap();
assert_eq!(event3, SessionEvent::StatusIdle { stop_reason: None });
assert!(event_stream.next().await.is_none());
}
#[tokio::test]
async fn test_unknown_event_type_produces_unknown_variant() {
let chunks = vec![
"event: some.future.event\ndata: {\"type\":\"some.future.event\",\"foo\":\"bar\"}\n\n",
];
let stream = byte_stream_from_chunks(chunks);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(event, SessionEvent::Unknown);
}
#[tokio::test]
async fn test_invalid_json_skips_event_and_continues() {
let chunks = vec![
"event: agent.message\ndata: {invalid json}\n\nevent: agent.message\ndata: {\"type\":\"agent.message\",\"content\":\"valid\"}\n\n",
];
let stream = byte_stream_from_chunks(chunks);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentMessage { id: None, content: serde_json::json!("valid") }
);
assert!(event_stream.next().await.is_none());
}
#[tokio::test]
async fn test_invalid_utf8_yields_encoding_error() {
let invalid_bytes: Vec<u8> = vec![0xFF, 0xFE, 0xFD];
let valid_chunk = "event: agent.message\ndata: {\"type\":\"agent.message\",\"content\":\"after error\"}\n\n";
let items: Vec<std::result::Result<Bytes, reqwest::Error>> =
vec![Ok(Bytes::from(invalid_bytes)), Ok(Bytes::from(valid_chunk.to_string()))];
let stream = stream::iter(items);
let mut event_stream = process_managed_agents_sse(Box::pin(stream), Duration::from_secs(5));
let result = event_stream.next().await.unwrap();
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(err, Error::Encoding { .. }));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentMessage { id: None, content: serde_json::json!("after error") }
);
}
#[tokio::test]
async fn test_timeout_yields_timeout_error() {
let stream = stream::pending::<std::result::Result<Bytes, reqwest::Error>>();
let mut event_stream =
process_managed_agents_sse(Box::pin(stream), Duration::from_millis(50));
let result = event_stream.next().await.unwrap();
assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(err, Error::Timeout { .. }));
}
#[tokio::test]
async fn test_chunked_data_across_multiple_reads() {
let chunks = vec![
"event: agent.message\n",
"data: {\"type\":\"agent.message\",",
"\"content\":\"split\"}\n\n",
];
let stream = byte_stream_from_chunks(chunks);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentMessage { id: None, content: serde_json::json!("split") }
);
}
#[tokio::test]
async fn test_tool_use_event_parsing() {
let data =
r#"{"type":"agent.tool_use","id":"tu_123","name":"bash","input":{"command":"ls"}}"#;
let chunk = format!("event: agent.tool_use\ndata: {data}\n\n");
let items: Vec<std::result::Result<Bytes, reqwest::Error>> = vec![Ok(Bytes::from(chunk))];
let stream = stream::iter(items);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentToolUse {
id: Some("tu_123".to_string()),
name: Some("bash".to_string()),
input: Some(serde_json::json!({"command": "ls"})),
}
);
}
#[tokio::test]
async fn test_custom_tool_use_event_parsing() {
let data = r#"{"type":"agent.custom_tool_use","id":"ctu_456","name":"my_tool","input":{"key":"value"}}"#;
let items: Vec<std::result::Result<Bytes, reqwest::Error>> =
vec![Ok(Bytes::from(format!("event: agent.custom_tool_use\ndata: {data}\n\n")))];
let stream = stream::iter(items);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentCustomToolUse {
id: Some("ctu_456".to_string()),
name: Some("my_tool".to_string()),
input: Some(serde_json::json!({"key": "value"})),
}
);
}
#[tokio::test]
async fn test_error_event_parsing() {
let data =
r#"{"type":"session.error","message":"something went wrong","code":"internal_error"}"#;
let items: Vec<std::result::Result<Bytes, reqwest::Error>> =
vec![Ok(Bytes::from(format!("event: session.error\ndata: {data}\n\n")))];
let stream = stream::iter(items);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::Error { error: None, message: Some("something went wrong".to_string()) }
);
}
#[tokio::test]
async fn test_empty_stream_produces_no_events() {
let stream = stream::empty::<std::result::Result<Bytes, reqwest::Error>>();
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
assert!(event_stream.next().await.is_none());
}
#[tokio::test]
async fn test_data_only_without_event_line() {
let chunks = vec!["data: {\"type\":\"agent.message\",\"content\":\"no event line\"}\n\n"];
let stream = byte_stream_from_chunks(chunks);
let mut event_stream = process_managed_agents_sse(stream, Duration::from_secs(5));
let event = event_stream.next().await.unwrap().unwrap();
assert_eq!(
event,
SessionEvent::AgentMessage { id: None, content: serde_json::json!("no event line") }
);
}
}