1use 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
12pub type BoxChatCompletionStream =
14 ChatCompletionStream<Pin<Box<dyn Stream<Item = Result<Bytes, reqwest::Error>> + Send>>>;
15
16pub 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 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 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 fn process_line(line: &str) -> (Option<Result<ChatCompletionChunk, ClientError>>, bool) {
65 let trimmed = line.trim();
66 if trimmed.is_empty() || trimmed.starts_with(':') {
67 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 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
112const 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 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 continue;
162 }
163
164 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 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 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); 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 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); 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 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}