Skip to main content

cera_client/
stream.rs

1//! Server-Sent Events (SSE) stream decoder for chat completion chunks.
2
3use std::pin::Pin;
4use std::task::{Context, Poll};
5
6use bytes::{Bytes, BytesMut};
7use futures_core::stream::{FusedStream, Stream};
8
9use crate::error::{ApiErrorEnvelope, ClientError};
10use crate::types::ChatCompletionChunk;
11
12/// A pinned, heap-allocated stream yielding incremental [`ChatCompletionChunk`] updates.
13pub type BoxChatCompletionStream =
14    ChatCompletionStream<Pin<Box<dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send>>>;
15
16/// Asynchronous stream yielding incremental [`ChatCompletionChunk`] updates from an SSE stream.
17pub struct ChatCompletionStream<S> {
18    inner: S,
19    buffer: BytesMut,
20    done: bool,
21}
22
23impl<S> std::fmt::Debug for ChatCompletionStream<S> {
24    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25        f.debug_struct("ChatCompletionStream")
26            .field("buffered_bytes", &self.buffer.len())
27            .field("done", &self.done)
28            .finish()
29    }
30}
31
32impl<S> ChatCompletionStream<S> {
33    /// Create a new stream wrapper around an inner byte stream.
34    pub fn new(inner: S) -> Self {
35        Self {
36            inner,
37            buffer: BytesMut::new(),
38            done: false,
39        }
40    }
41}
42
43impl<S> ChatCompletionStream<S>
44where
45    S: Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
46{
47    /// Pin and box this stream into a [`BoxChatCompletionStream`].
48    pub fn boxed(self) -> BoxChatCompletionStream {
49        ChatCompletionStream {
50            inner: Box::pin(self.inner),
51            buffer: self.buffer,
52            done: self.done,
53        }
54    }
55}
56
57impl<S> ChatCompletionStream<S> {
58    /// Helper to process a single decoded line.
59    ///
60    /// Returns:
61    /// - `Some(Ok(chunk))` if a valid chunk was parsed
62    /// - `Some(Err(err))` if an error or in-band API error occurred
63    /// - `None` if the line was empty, a comment, other SSE field, or `[DONE]`
64    fn process_line(line: &str) -> (Option<Result<ChatCompletionChunk, ClientError>>, bool) {
65        let trimmed = line.trim();
66        if trimmed.is_empty() || trimmed.starts_with(':') {
67            // Keep-alive or comment line, ignore.
68            return (None, false);
69        }
70
71        if let Some(payload) = trimmed.strip_prefix("data:") {
72            let data = payload.trim();
73            if data.is_empty() {
74                // Keep-alive empty data heartbeat, ignore.
75                return (None, false);
76            }
77            if data == "[DONE]" {
78                return (None, true);
79            }
80
81            match serde_json::from_str::<ChatCompletionChunk>(data) {
82                Ok(chunk) => (Some(Ok(chunk)), false),
83                Err(err) => {
84                    if let Ok(env) = serde_json::from_str::<ApiErrorEnvelope>(data) {
85                        (
86                            Some(Err(ClientError::Api {
87                                status: None,
88                                message: env.error.message,
89                                error_type: env.error.error_type,
90                                code: env.error.code,
91                                param: env.error.param,
92                            })),
93                            true,
94                        )
95                    } else {
96                        (
97                            Some(Err(ClientError::Serialization {
98                                source: err,
99                                raw_payload: Some(data.to_string()),
100                            })),
101                            false,
102                        )
103                    }
104                }
105            }
106        } else {
107            (None, false)
108        }
109    }
110}
111
112/// Maximum allowed buffer capacity for an SSE line (16 MB) to guard against unbounded streams.
113const MAX_STREAM_BUFFER_BYTES: usize = 16 * 1024 * 1024;
114
115impl<S> Stream for ChatCompletionStream<S>
116where
117    S: Stream<Item = Result<Bytes, reqwest::Error>> + Unpin,
118{
119    type Item = Result<ChatCompletionChunk, ClientError>;
120
121    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
122        let this = self.as_mut().get_mut();
123
124        loop {
125            if this.done {
126                return Poll::Ready(None);
127            }
128
129            // Check if we have a full line in the buffer
130            if let Some(idx) = this.buffer.iter().position(|&b| b == b'\n') {
131                let line_bytes = this.buffer.split_to(idx + 1);
132                let mut slice = line_bytes.as_ref();
133                if slice.ends_with(b"\n") {
134                    slice = &slice[..slice.len() - 1];
135                }
136                if slice.ends_with(b"\r") {
137                    slice = &slice[..slice.len() - 1];
138                }
139
140                let line = match std::str::from_utf8(slice) {
141                    Ok(s) => s,
142                    Err(e) => {
143                        this.done = true;
144                        return Poll::Ready(Some(Err(ClientError::Stream(format!(
145                            "Invalid UTF-8 in SSE stream line: {e}"
146                        )))));
147                    }
148                };
149
150                let (result, is_done) = Self::process_line(line);
151                if is_done {
152                    this.done = true;
153                }
154                if let Some(res) = result {
155                    return Poll::Ready(Some(res));
156                }
157                if is_done {
158                    return Poll::Ready(None);
159                }
160                // If it was a comment or empty line, continue processing buffer
161                continue;
162            }
163
164            // Need more data from the network
165            match Pin::new(&mut this.inner).poll_next(cx) {
166                Poll::Ready(Some(Ok(bytes))) => {
167                    if this.buffer.len() + bytes.len() > MAX_STREAM_BUFFER_BYTES {
168                        this.done = true;
169                        return Poll::Ready(Some(Err(ClientError::Stream(
170                            "SSE stream line exceeded maximum buffer capacity of 16 MB".to_string(),
171                        ))));
172                    }
173                    this.buffer.extend_from_slice(&bytes);
174                }
175                Poll::Ready(Some(Err(e))) => {
176                    this.done = true;
177                    return Poll::Ready(Some(Err(ClientError::Http(e))));
178                }
179                Poll::Ready(None) => {
180                    // Stream closed
181                    if !this.buffer.is_empty() {
182                        let remaining = std::mem::take(&mut this.buffer);
183                        let mut slice = remaining.as_ref();
184                        if slice.ends_with(b"\n") {
185                            slice = &slice[..slice.len() - 1];
186                        }
187                        if slice.ends_with(b"\r") {
188                            slice = &slice[..slice.len() - 1];
189                        }
190
191                        if !slice.is_empty() {
192                            let line = match std::str::from_utf8(slice) {
193                                Ok(s) => s,
194                                Err(e) => {
195                                    this.done = true;
196                                    return Poll::Ready(Some(Err(ClientError::Stream(format!(
197                                        "Invalid UTF-8 in trailing SSE data: {e}"
198                                    )))));
199                                }
200                            };
201
202                            let (result, is_done) = Self::process_line(line);
203                            this.done = true;
204                            if let Some(res) = result {
205                                return Poll::Ready(Some(res));
206                            }
207                            if is_done {
208                                return Poll::Ready(None);
209                            }
210                        }
211                    }
212                    this.done = true;
213                    return Poll::Ready(None);
214                }
215                Poll::Pending => return Poll::Pending,
216            }
217        }
218    }
219}
220
221impl<S> FusedStream for ChatCompletionStream<S>
222where
223    S: Stream<Item = Result<Bytes, reqwest::Error>> + Unpin,
224{
225    fn is_terminated(&self) -> bool {
226        self.done
227    }
228}
229
230#[cfg(test)]
231mod tests {
232    use super::*;
233    use futures_util::StreamExt;
234
235    #[tokio::test]
236    async fn test_sse_stream_parsing_with_fragmented_chunks() {
237        let chunk1 = Bytes::from(
238            "data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"created\":123,\"model\":\"m\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hel",
239        );
240        let chunk2 = Bytes::from(
241            "lo\"},\"finish_reason\":null}]}\n\n: keep-alive\n\ndata: {\"id\":\"2\",\"object\":\"chat.completion.chunk\",\"created\":124,\"model\":\"m\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n",
242        );
243
244        let byte_stream = futures_util::stream::iter(vec![
245            Ok::<Bytes, reqwest::Error>(chunk1),
246            Ok::<Bytes, reqwest::Error>(chunk2),
247        ]);
248
249        let mut sse_stream = ChatCompletionStream::new(byte_stream);
250
251        let item1 = sse_stream.next().await.expect("item 1").unwrap();
252        assert_eq!(item1.id, "1");
253        assert_eq!(item1.choices[0].delta.content.as_deref(), Some("Hello"));
254
255        let item2 = sse_stream.next().await.expect("item 2").unwrap();
256        assert_eq!(item2.id, "2");
257        assert_eq!(item2.choices[0].delta.content.as_deref(), Some(" world"));
258        assert_eq!(item2.choices[0].finish_reason.as_deref(), Some("stop"));
259
260        assert!(sse_stream.next().await.is_none());
261    }
262
263    #[tokio::test]
264    async fn test_sse_stream_handles_multibyte_utf8_split_across_chunks() {
265        // Japanese character "あ" is [0xE3, 0x81, 0x82] in UTF-8
266        let prefix = "data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"m\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"";
267        let suffix = "\"},\"finish_reason\":null}]}\n\ndata: [DONE]\n\n";
268
269        let mut chunk1_bytes = prefix.as_bytes().to_vec();
270        chunk1_bytes.push(0xE3);
271        chunk1_bytes.push(0x81); // Split before last byte
272
273        let mut chunk2_bytes = vec![0x82];
274        chunk2_bytes.extend_from_slice(suffix.as_bytes());
275
276        let byte_stream = futures_util::stream::iter(vec![
277            Ok::<Bytes, reqwest::Error>(Bytes::from(chunk1_bytes)),
278            Ok::<Bytes, reqwest::Error>(Bytes::from(chunk2_bytes)),
279        ]);
280
281        let mut sse_stream = ChatCompletionStream::new(byte_stream);
282        let item = sse_stream.next().await.expect("item 1").unwrap();
283        assert_eq!(item.choices[0].delta.content.as_deref(), Some("あ"));
284        assert!(sse_stream.next().await.is_none());
285    }
286
287    #[tokio::test]
288    async fn test_sse_stream_handles_four_byte_emoji_split_across_chunks() {
289        // "🦀" is UTF-8: [0xF0, 0x9F, 0xA6, 0x80]
290        let prefix = "data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"m\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"";
291        let suffix = "\"},\"finish_reason\":null}]}\n\ndata: [DONE]\n\n";
292
293        let mut chunk1_bytes = prefix.as_bytes().to_vec();
294        chunk1_bytes.push(0xF0);
295        chunk1_bytes.push(0x9F); // Split mid 4-byte sequence
296
297        let mut chunk2_bytes = vec![0xA6, 0x80];
298        chunk2_bytes.extend_from_slice(suffix.as_bytes());
299
300        let byte_stream = futures_util::stream::iter(vec![
301            Ok::<Bytes, reqwest::Error>(Bytes::from(chunk1_bytes)),
302            Ok::<Bytes, reqwest::Error>(Bytes::from(chunk2_bytes)),
303        ]);
304
305        let mut sse_stream = ChatCompletionStream::new(byte_stream);
306        let item = sse_stream.next().await.expect("item 1").unwrap();
307        assert_eq!(item.choices[0].delta.content.as_deref(), Some("🦀"));
308        assert!(sse_stream.next().await.is_none());
309    }
310
311    #[tokio::test]
312    async fn test_sse_stream_handles_midstream_api_error() {
313        let payload = "data: {\"error\":{\"message\":\"Model overloaded\",\"type\":\"server_error\",\"code\":\"server_error\"}}\n\n";
314        let byte_stream =
315            futures_util::stream::iter(vec![Ok::<Bytes, reqwest::Error>(Bytes::from(payload))]);
316        let mut sse_stream = ChatCompletionStream::new(byte_stream);
317
318        let err = sse_stream.next().await.expect("item").unwrap_err();
319        match err {
320            ClientError::Api { message, code, .. } => {
321                assert_eq!(message, "Model overloaded");
322                assert_eq!(code.as_deref(), Some("server_error"));
323            }
324            other => panic!("expected Api error, got: {other:?}"),
325        }
326        assert!(sse_stream.next().await.is_none());
327    }
328
329    #[tokio::test]
330    async fn test_sse_stream_handles_bare_string_api_error() {
331        let payload = "data: {\"error\":\"rate limit reached, please slow down\"}\n\n";
332        let byte_stream =
333            futures_util::stream::iter(vec![Ok::<Bytes, reqwest::Error>(Bytes::from(payload))]);
334        let mut sse_stream = ChatCompletionStream::new(byte_stream);
335
336        let err = sse_stream.next().await.expect("item").unwrap_err();
337        match err {
338            ClientError::Api { message, .. } => {
339                assert_eq!(message, "rate limit reached, please slow down");
340            }
341            other => panic!("expected Api error, got: {other:?}"),
342        }
343        assert!(sse_stream.next().await.is_none());
344    }
345
346    #[tokio::test]
347    async fn test_sse_stream_handles_invalid_utf8() {
348        let bad_bytes = Bytes::from(vec![0xff, 0xfe, 0xfd, b'\n']);
349        let byte_stream = futures_util::stream::iter(vec![Ok::<Bytes, reqwest::Error>(bad_bytes)]);
350        let mut sse_stream = ChatCompletionStream::new(byte_stream);
351
352        let err = sse_stream.next().await.expect("error").unwrap_err();
353        match err {
354            ClientError::Stream(msg) => assert!(msg.contains("Invalid UTF-8")),
355            other => panic!("unexpected error: {other:?}"),
356        }
357    }
358
359    #[tokio::test]
360    async fn test_sse_stream_handles_empty_data_heartbeat() {
361        let stream_text = "data: \n\ndata: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"m\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}\n\ndata:\n\ndata: [DONE]\n\n";
362        let byte_stream =
363            futures_util::stream::iter(vec![Ok::<Bytes, reqwest::Error>(Bytes::from(stream_text))]);
364        let mut sse_stream = ChatCompletionStream::new(byte_stream);
365
366        let item = sse_stream.next().await.expect("item").unwrap();
367        assert_eq!(item.choices[0].delta.content.as_deref(), Some("ok"));
368        assert!(sse_stream.next().await.is_none());
369    }
370
371    #[tokio::test]
372    async fn test_fused_stream_and_boxed() {
373        let stream_text = "data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"m\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"test\"},\"finish_reason\":\"stop\"}]}\n\ndata: [DONE]\n\n";
374        let byte_stream =
375            futures_util::stream::iter(vec![Ok::<Bytes, reqwest::Error>(Bytes::from(stream_text))]);
376        let mut boxed_stream = ChatCompletionStream::new(byte_stream).boxed();
377
378        assert!(!boxed_stream.is_terminated());
379        let item = boxed_stream.next().await.expect("item").unwrap();
380        assert_eq!(item.choices[0].delta.content.as_deref(), Some("test"));
381        assert!(boxed_stream.next().await.is_none());
382        assert!(boxed_stream.is_terminated());
383        // FusedStream guarantees subsequent polls return None without panicking
384        assert!(boxed_stream.next().await.is_none());
385    }
386
387    #[tokio::test]
388    async fn test_sse_stream_handles_serialization_error_with_raw_payload() {
389        let payload = "data: {\"not_valid_json_chunk\": true}\n\n";
390        let byte_stream =
391            futures_util::stream::iter(vec![Ok::<Bytes, reqwest::Error>(Bytes::from(payload))]);
392        let mut sse_stream = ChatCompletionStream::new(byte_stream);
393
394        let err = sse_stream.next().await.expect("item").unwrap_err();
395        match err {
396            ClientError::Serialization {
397                source: _,
398                raw_payload,
399            } => {
400                assert_eq!(
401                    raw_payload.as_deref(),
402                    Some("{\"not_valid_json_chunk\": true}")
403                );
404            }
405            other => panic!("expected Serialization error, got: {other:?}"),
406        }
407    }
408}