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>,
}
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(),
}
}
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()))
}
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'],
};
let result = reader.take_line();
assert!(result.is_err(), "invalid UTF-8 must surface as an error");
}
#[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(),
};
let line = reader.take_line().unwrap().unwrap();
assert_eq!(line, "data: hello");
}
}