Skip to main content

gglib_core/sse/
decoder.rs

1//! Stateful SSE byte-stream decoder.
2//!
3//! [`SseStreamDecoder`] accumulates raw bytes from an HTTP response into a line
4//! buffer, drains complete `data:` lines, and delegates frame parsing to
5//! [`super::parser`].  Its explicit state makes it straightforward to unit-
6//! test without standing up an actual HTTP server or wrapping everything in an
7//! `async_stream` macro block.
8
9use anyhow::Result;
10use tracing::debug;
11
12use crate::LlmStreamEvent;
13
14use super::parser::{SseParseResult, parse_sse_frame};
15
16/// Stateful decoder that turns a sequence of raw SSE byte chunks into a
17/// sequence of [`LlmStreamEvent`] values.
18///
19/// # Usage
20///
21/// ```ignore
22/// let mut decoder = SseStreamDecoder::default();
23/// while let Some(chunk) = byte_stream.next().await { … }
24///     let (events, stop) = decoder.feed_bytes(&chunk);
25///     for event in events { … }
26///     if stop { break; }
27/// }
28/// if let Some(fallback) = decoder.finish() { … }
29/// ```
30#[derive(Default)]
31pub struct SseStreamDecoder {
32    buf: String,
33    /// Set to `true` once a [`LlmStreamEvent::Done`] has been yielded, so the
34    /// `[DONE]` sentinel doesn't generate a duplicate.
35    done_sent: bool,
36}
37
38impl SseStreamDecoder {
39    /// Feed one raw byte chunk into the decoder.
40    ///
41    /// Returns `(events, should_stop)`.
42    ///
43    /// - `events` — zero or more parsed [`LlmStreamEvent`] values (or stream
44    ///   errors) extracted from the bytes fed so far.
45    /// - `should_stop` — `true` when the SSE stream has reached its natural end
46    ///   (a `[DONE]` sentinel or an unrecoverable parse error).  The caller
47    ///   must not feed any further chunks once this flag is `true`.
48    pub fn feed_bytes(&mut self, bytes: &[u8]) -> (Vec<Result<LlmStreamEvent>>, bool) {
49        let text = match std::str::from_utf8(bytes) {
50            Ok(t) => t,
51            Err(e) => {
52                return (
53                    vec![Err(anyhow::anyhow!("invalid UTF-8 in LLM SSE stream: {e}"))],
54                    true,
55                );
56            }
57        };
58        self.buf.push_str(text);
59        let mut events = Vec::new();
60
61        while let Some(newline_pos) = self.buf.find('\n') {
62            let line = self.buf[..newline_pos].trim_end_matches('\r').to_owned();
63            self.buf.drain(..=newline_pos);
64
65            // Skip blank lines and SSE comment lines.
66            let Some(data) = line.strip_prefix("data: ") else {
67                continue;
68            };
69
70            match parse_sse_frame(data) {
71                Ok(SseParseResult::Done) => {
72                    if !self.done_sent {
73                        debug!(
74                            "LLM stream ended with [DONE] but no prior finish_reason \
75                             — emitting fallback Done"
76                        );
77                        events.push(Ok(LlmStreamEvent::Done {
78                            finish_reason: "stop".to_owned(),
79                        }));
80                    }
81                    self.done_sent = true;
82                    return (events, true);
83                }
84                Ok(SseParseResult::Events(parsed_events)) => {
85                    let mut saw_terminal_error = false;
86                    for event in parsed_events {
87                        if matches!(event, LlmStreamEvent::Done { .. }) {
88                            self.done_sent = true;
89                        }
90                        if matches!(event, LlmStreamEvent::UpstreamError { .. }) {
91                            // Terminal condition: the encoder appends its own
92                            // `[DONE]` sentinel right after this event (see
93                            // `SseEncoder::encode`), and nothing meaningful
94                            // is expected to follow an inline upstream
95                            // error. Stop feeding further bytes, same as the
96                            // literal `[DONE]` sentinel case above, so a
97                            // stray fallback `Done` isn't appended by
98                            // `finish()`.
99                            saw_terminal_error = true;
100                            self.done_sent = true;
101                        }
102                        events.push(Ok(event));
103                    }
104                    if saw_terminal_error {
105                        return (events, true);
106                    }
107                }
108                Err(e) => {
109                    events.push(Err(e));
110                    return (events, true);
111                }
112            }
113        }
114
115        (events, false)
116    }
117
118    /// Emit a fallback `Done` event if the byte stream ended without one.
119    ///
120    /// Call this once after the upstream byte stream is fully exhausted.
121    /// Returns `None` if a `Done` was already yielded by [`feed_bytes`].
122    #[must_use]
123    pub fn finish(self) -> Option<LlmStreamEvent> {
124        if self.done_sent {
125            None
126        } else {
127            debug!("LLM byte-stream ended without [DONE] sentinel — emitting fallback Done");
128            Some(LlmStreamEvent::Done {
129                finish_reason: "stop".to_owned(),
130            })
131        }
132    }
133}
134
135// =============================================================================
136// Tests
137// =============================================================================
138
139#[cfg(test)]
140mod tests {
141    use anyhow::Result;
142
143    use super::SseStreamDecoder;
144    use crate::LlmStreamEvent;
145
146    fn text_delta_frame(text: &str) -> String {
147        let json = serde_json::json!({
148            "choices": [{
149                "delta": { "content": text },
150                "finish_reason": null
151            }]
152        });
153        format!("data: {json}\n")
154    }
155
156    fn done_frame() -> &'static str {
157        "data: [DONE]\n"
158    }
159
160    fn finish_reason_frame() -> String {
161        let json = serde_json::json!({
162            "choices": [{
163                "delta": {},
164                "finish_reason": "stop"
165            }]
166        });
167        format!("data: {json}\n")
168    }
169
170    // ---- helpers ------------------------------------------------------------
171
172    fn collect_all(decoder: &mut SseStreamDecoder, input: &str) -> (Vec<LlmStreamEvent>, bool) {
173        let (raw, stop) = decoder.feed_bytes(input.as_bytes());
174        let events: Vec<_> = raw.into_iter().map(Result::unwrap).collect();
175        (events, stop)
176    }
177
178    // ---- tests --------------------------------------------------------------
179
180    #[test]
181    fn text_delta_is_emitted() {
182        let mut dec = SseStreamDecoder::default();
183        let (events, stop) = collect_all(&mut dec, &text_delta_frame("hello"));
184        assert!(!stop);
185        assert!(
186            events
187                .iter()
188                .any(|e| matches!(e, LlmStreamEvent::TextDelta { content } if content == "hello"))
189        );
190    }
191
192    #[test]
193    fn done_sentinel_signals_stop_and_emits_fallback() {
194        let mut dec = SseStreamDecoder::default();
195        let (events, stop) = collect_all(&mut dec, done_frame());
196        assert!(stop, "decoder should signal stop on [DONE]");
197        assert!(
198            events
199                .iter()
200                .any(|e| matches!(e, LlmStreamEvent::Done { .. })),
201            "fallback Done should be emitted when no prior finish_reason"
202        );
203        assert!(
204            dec.finish().is_none(),
205            "finish() must return None after a [DONE] sentinel — done_sent must be set"
206        );
207    }
208
209    #[test]
210    fn finish_reason_then_done_no_duplicate_done() {
211        let mut dec = SseStreamDecoder::default();
212        let input = format!("{}{}", finish_reason_frame(), done_frame());
213        let (events, stop) = collect_all(&mut dec, &input);
214        assert!(stop);
215        let done_count = events
216            .iter()
217            .filter(|e| matches!(e, LlmStreamEvent::Done { .. }))
218            .count();
219        assert_eq!(done_count, 1, "exactly one Done should be emitted");
220    }
221
222    #[test]
223    fn finish_emits_fallback_when_stream_ends_without_done() {
224        let mut dec = SseStreamDecoder::default();
225        let _ = collect_all(&mut dec, &text_delta_frame("partial"));
226        let fallback = dec.finish();
227        assert!(
228            fallback.is_some(),
229            "finish() should return a fallback Done when stream ends without one"
230        );
231    }
232
233    #[test]
234    fn finish_returns_none_when_done_already_sent() {
235        let mut dec = SseStreamDecoder::default();
236        let _ = collect_all(&mut dec, &finish_reason_frame());
237        assert!(
238            dec.finish().is_none(),
239            "finish() must not emit a second Done"
240        );
241    }
242
243    #[test]
244    fn inline_error_frame_signals_stop_and_suppresses_fallback_done() {
245        let mut dec = SseStreamDecoder::default();
246        let frame = format!(
247            "data: {}\n",
248            serde_json::json!({ "error": { "message": "boom" } })
249        );
250        let (events, stop) = collect_all(&mut dec, &frame);
251        assert!(stop, "inline error frame should signal stop");
252        assert_eq!(events.len(), 1);
253        assert!(matches!(&events[0], LlmStreamEvent::UpstreamError { .. }));
254        assert!(
255            dec.finish().is_none(),
256            "finish() must not append a fallback Done after an inline error"
257        );
258    }
259
260    #[test]
261    fn partial_line_buffered_until_newline_arrives() {
262        let mut dec = SseStreamDecoder::default();
263        let full_frame = text_delta_frame("world");
264
265        let mid = full_frame.len() / 2;
266        let (first_events, stop1) = collect_all(&mut dec, &full_frame[..mid]);
267        assert!(!stop1);
268        assert!(first_events.is_empty(), "no complete line yet");
269
270        let (second_events, stop2) = collect_all(&mut dec, &full_frame[mid..]);
271        assert!(!stop2);
272        assert!(
273            second_events
274                .iter()
275                .any(|e| matches!(e, LlmStreamEvent::TextDelta { .. })),
276            "TextDelta should be emitted once the newline arrives"
277        );
278    }
279}