use std::pin::Pin;
use futures::stream::{Stream, StreamExt};
use reqwest::Response;
use crate::api::error::ApiError;
const SSE_MAX_BUFFER: usize = 1024 * 1024;
pub(super) struct SseReader {
pub(super) bytes: Pin<Box<dyn Stream<Item = Result<bytes::Bytes, ApiError>> + Send>>,
pub(super) buf: Vec<u8>,
#[cfg(feature = "openai")]
pub(super) done_marker_seen: bool,
}
impl SseReader {
pub(super) fn from_response(resp: Response) -> Self {
let bytes = resp
.bytes_stream()
.map(|res| res.map_err(|e| ApiError::http(e.to_string())));
Self {
bytes: Box::pin(bytes),
buf: Vec::new(),
#[cfg(feature = "openai")]
done_marker_seen: false,
}
}
pub(super) fn take_line(&mut self) -> Result<Option<String>, ApiError> {
let Some(pos) = self.buf.iter().position(|&b| b == b'\n') else {
return Ok(None);
};
let rest_start = pos.saturating_add(1);
let line_bytes: Vec<u8> = self.buf.drain(..rest_start).collect();
let line = String::from_utf8(line_bytes)
.map_err(|e| ApiError::http(format!("SSE line is not valid UTF-8: {e}")))?;
Ok(Some(line.trim().to_string()))
}
#[cfg(feature = "openai")]
pub(super) fn done_marker_seen(&self) -> bool {
self.done_marker_seen
}
#[cfg(feature = "openai")]
pub(super) fn mark_done_marker_seen(&mut self) {
self.done_marker_seen = true;
}
pub(super) async fn next_chunk(&mut self) -> Result<Option<()>, ApiError> {
match self.bytes.next().await {
Some(Ok(chunk)) => {
self.buf.extend_from_slice(&chunk);
if self.buf.len() > SSE_MAX_BUFFER {
return Err(ApiError::http(format!(
"SSE buffer exceeded {SSE_MAX_BUFFER} bytes"
)));
}
Ok(Some(()))
}
Some(Err(e)) => Err(e),
None => Ok(None),
}
}
}
#[cfg(any(feature = "openai", feature = "anthropic", feature = "gemini"))]
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn take_line_invalid_utf8_returns_error() {
let mut reader = SseReader {
bytes: Box::pin(futures::stream::empty()),
buf: vec![0xFF, 0xFE, 0xFD, b'\n'],
#[cfg(feature = "openai")]
done_marker_seen: false,
};
let result = reader.take_line();
assert!(result.is_err(), "invalid UTF-8 must surface as an error");
}
#[tokio::test]
async fn buffer_overflow_errors_instead_of_growing_unbounded() {
let big = vec![b'a'; super::SSE_MAX_BUFFER + 1];
let mut reader = SseReader {
bytes: Box::pin(futures::stream::iter(vec![Ok(bytes::Bytes::from(big))])),
buf: Vec::new(),
#[cfg(feature = "openai")]
done_marker_seen: false,
};
let err = reader
.next_chunk()
.await
.expect_err("a newline-less oversize chunk must error, not buffer");
assert!(
err.to_string().contains("exceeded"),
"the error names the guard: {err}"
);
}
#[test]
fn take_line_valid_utf8_returns_ok() {
let mut reader = SseReader {
bytes: Box::pin(futures::stream::empty()),
buf: b"data: hello\n".to_vec(),
#[cfg(feature = "openai")]
done_marker_seen: false,
};
let line = reader.take_line().unwrap().unwrap();
assert_eq!(line, "data: hello");
}
}