Skip to main content

llmrix_rust_sdk/streaming/
parser.rs

1use bytes::Bytes;
2use futures_util::{Stream, StreamExt};
3use serde_json::Value;
4
5use super::event::*;
6use crate::error::{LlmrixError, Result};
7
8/// Callback invoked once per parsed [`StreamEvent`].
9/// Return `Err` to abort streaming early.
10pub type EventHandler = Box<dyn FnMut(&StreamEvent) -> Result<()> + Send>;
11
12// ---------------------------------------------------------------------------
13// Registry  (channel:type → deserializer)
14// ---------------------------------------------------------------------------
15
16fn dispatch_frame(event_name: &str, data: &str) -> Option<StreamEvent> {
17    // Bare keep-alive
18    if data == ":ping" || event_name == "heartbeat" {
19        return Some(StreamEvent::Heartbeat);
20    }
21
22    let raw: Value = serde_json::from_str(data).ok()?;
23    let channel = raw["channel"].as_str().unwrap_or(event_name);
24    let typ = raw["type"].as_str().unwrap_or("");
25
26    let evt = match (channel, typ) {
27        ("lifecycle", "run_start")    => StreamEvent::RunStart(serde_json::from_value(raw).ok()?),
28        ("lifecycle", "run_end")      => StreamEvent::RunEnd(serde_json::from_value(raw).ok()?),
29        ("messages",  "message_chunk")=> StreamEvent::MessageChunk(serde_json::from_value(raw).ok()?),
30        ("tools",     "tool_start")   => StreamEvent::ToolStart(serde_json::from_value(raw).ok()?),
31        ("tools",     "tool_end")     => StreamEvent::ToolEnd(serde_json::from_value(raw).ok()?),
32        ("tools",     "subagent_start")=>StreamEvent::SubagentStart(serde_json::from_value(raw).ok()?),
33        ("tools",     "subagent_end") => StreamEvent::SubagentEnd(serde_json::from_value(raw).ok()?),
34        ("hitl",      "hitl_interrupt")=>StreamEvent::HitlInterrupt(serde_json::from_value(raw).ok()?),
35        ("error",     "error")        => StreamEvent::Error(serde_json::from_value(raw).ok()?),
36        ("error",     "cancelled")    => StreamEvent::Cancelled(serde_json::from_value(raw).ok()?),
37        ("heartbeat", _)              => StreamEvent::Heartbeat,
38        _ => return None, // unknown — skip
39    };
40    Some(evt)
41}
42
43// ---------------------------------------------------------------------------
44// Async SSE parser
45// ---------------------------------------------------------------------------
46
47/// Reads a byte stream from a live SSE response, parses frames, and dispatches
48/// each event to `handler`. Blocks (awaits) until the server closes the stream.
49pub(crate) async fn parse_sse_stream<S, F>(
50    bytes_stream: S,
51    mut handler: F,
52) -> Result<()>
53where
54    S: Stream<Item = reqwest::Result<Bytes>> + Unpin,
55    F: FnMut(&StreamEvent) -> Result<()>,
56{
57    let mut bytes_stream = bytes_stream;
58    let mut remainder: Vec<u8> = Vec::new();
59    let mut event_name = String::new();
60    let mut data_lines: Vec<String> = Vec::new();
61
62    while let Some(chunk) = bytes_stream.next().await {
63        let chunk = chunk.map_err(LlmrixError::Transport)?;
64        remainder.extend_from_slice(&chunk);
65
66        // Process all complete newline-terminated lines in the buffer.
67        loop {
68            match remainder.iter().position(|&b| b == b'\n') {
69                None => break,
70                Some(pos) => {
71                    // Trim trailing \r
72                    let end = if pos > 0 && remainder[pos - 1] == b'\r' { pos - 1 } else { pos };
73                    let line = String::from_utf8_lossy(&remainder[..end]).into_owned();
74                    remainder = remainder[pos + 1..].to_vec();
75
76                    if line.is_empty() {
77                        // Blank line → dispatch accumulated frame
78                        if !data_lines.is_empty() {
79                            let data = data_lines.join("\n");
80                            data_lines.clear();
81                            if let Some(event) = dispatch_frame(&event_name, &data) {
82                                handler(&event)?;
83                            }
84                            event_name.clear();
85                        }
86                    } else if let Some(rest) = line.strip_prefix("event:") {
87                        event_name = rest.trim().to_string();
88                    } else if let Some(rest) = line.strip_prefix("data:") {
89                        let value = rest.strip_prefix(' ').unwrap_or(rest);
90                        data_lines.push(value.to_string());
91                    }
92                    // Ignore "id:" and "retry:" fields
93                }
94            }
95        }
96    }
97
98    // Flush any unterminated final frame
99    if !data_lines.is_empty() {
100        let data = data_lines.join("\n");
101        if let Some(event) = dispatch_frame(&event_name, &data) {
102            handler(&event)?;
103        }
104    }
105
106    Ok(())
107}