gglib_core/sse/
decoder.rs1use anyhow::Result;
10use tracing::debug;
11
12use crate::LlmStreamEvent;
13
14use super::parser::{SseParseResult, parse_sse_frame};
15
16#[derive(Default)]
31pub struct SseStreamDecoder {
32 buf: String,
33 done_sent: bool,
36}
37
38impl SseStreamDecoder {
39 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 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 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 #[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#[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 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 #[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}