use bytes::Bytes;
use futures_util::{Stream, StreamExt};
use serde_json::Value;
use super::event::*;
use crate::error::{LlmrixError, Result};
pub type EventHandler = Box<dyn FnMut(&StreamEvent) -> Result<()> + Send>;
fn dispatch_frame(event_name: &str, data: &str) -> Option<StreamEvent> {
if data == ":ping" || event_name == "heartbeat" {
return Some(StreamEvent::Heartbeat);
}
let raw: Value = serde_json::from_str(data).ok()?;
let channel = raw["channel"].as_str().unwrap_or(event_name);
let typ = raw["type"].as_str().unwrap_or("");
let evt = match (channel, typ) {
("lifecycle", "run_start") => StreamEvent::RunStart(serde_json::from_value(raw).ok()?),
("lifecycle", "run_end") => StreamEvent::RunEnd(serde_json::from_value(raw).ok()?),
("messages", "message_chunk")=> StreamEvent::MessageChunk(serde_json::from_value(raw).ok()?),
("tools", "tool_start") => StreamEvent::ToolStart(serde_json::from_value(raw).ok()?),
("tools", "tool_end") => StreamEvent::ToolEnd(serde_json::from_value(raw).ok()?),
("tools", "subagent_start")=>StreamEvent::SubagentStart(serde_json::from_value(raw).ok()?),
("tools", "subagent_end") => StreamEvent::SubagentEnd(serde_json::from_value(raw).ok()?),
("hitl", "hitl_interrupt")=>StreamEvent::HitlInterrupt(serde_json::from_value(raw).ok()?),
("error", "error") => StreamEvent::Error(serde_json::from_value(raw).ok()?),
("error", "cancelled") => StreamEvent::Cancelled(serde_json::from_value(raw).ok()?),
("heartbeat", _) => StreamEvent::Heartbeat,
_ => return None, };
Some(evt)
}
pub(crate) async fn parse_sse_stream<S, F>(
bytes_stream: S,
mut handler: F,
) -> Result<()>
where
S: Stream<Item = reqwest::Result<Bytes>> + Unpin,
F: FnMut(&StreamEvent) -> Result<()>,
{
let mut bytes_stream = bytes_stream;
let mut remainder: Vec<u8> = Vec::new();
let mut event_name = String::new();
let mut data_lines: Vec<String> = Vec::new();
while let Some(chunk) = bytes_stream.next().await {
let chunk = chunk.map_err(LlmrixError::Transport)?;
remainder.extend_from_slice(&chunk);
loop {
match remainder.iter().position(|&b| b == b'\n') {
None => break,
Some(pos) => {
let end = if pos > 0 && remainder[pos - 1] == b'\r' { pos - 1 } else { pos };
let line = String::from_utf8_lossy(&remainder[..end]).into_owned();
remainder = remainder[pos + 1..].to_vec();
if line.is_empty() {
if !data_lines.is_empty() {
let data = data_lines.join("\n");
data_lines.clear();
if let Some(event) = dispatch_frame(&event_name, &data) {
handler(&event)?;
}
event_name.clear();
}
} else if let Some(rest) = line.strip_prefix("event:") {
event_name = rest.trim().to_string();
} else if let Some(rest) = line.strip_prefix("data:") {
let value = rest.strip_prefix(' ').unwrap_or(rest);
data_lines.push(value.to_string());
}
}
}
}
}
if !data_lines.is_empty() {
let data = data_lines.join("\n");
if let Some(event) = dispatch_frame(&event_name, &data) {
handler(&event)?;
}
}
Ok(())
}