//! A server-sent-events reader: turns a streamed body into `data:` payloads. //! //! `llama-server` streams a completion as SSE — one JSON object per `data:` line, then //! `data: [DONE]`. Lines end in `\n` or `\r\n`. Blank lines, comment lines (`:`) and any field //! other than `data` carry nothing and are skipped. After `data:` one leading space is dropped, and //! only one. Every byte the peer sends is read in bounded chunks; a line is only UTF-8 decoded once //! its whole length has arrived. use std::io::{ErrorKind, Read}; use crate::http::is_timeout; #[derive(Debug, Clone, PartialEq, Eq)] pub enum SseItem { /// The text after `data:`, without the line ending. Data(String), Done, } #[derive(Debug)] pub enum SseError { Io(std::io::Error), Timeout, Truncated, LineTooLong, NotUtf8, } impl std::fmt::Display for SseError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { SseError::Io(e) => write!(f, "io error: {e}"), SseError::Timeout => write!(f, "read timed out before the line arrived"), SseError::Truncated => write!(f, "the body ended in the middle of a line"), SseError::LineTooLong => write!(f, "a line exceeded the limit"), SseError::NotUtf8 => write!(f, "a line was not valid UTF-8"), } } } impl std::error::Error for SseError { fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { match self { SseError::Io(e) => Some(e), _ => None, } } } pub struct Events { reader: R, buf: Vec, max_line: usize, ended: bool, } impl Events { pub fn new(reader: R, max_line: usize) -> Self { Events { reader, buf: Vec::new(), max_line, ended: false, } } /// The next item, or `None` when the stream has ended cleanly between lines. pub fn next_item(&mut self) -> Result, SseError> { loop { if self.ended { return Ok(None); } let line = match self.read_line() { Ok(None) => { self.ended = true; return Ok(None); } Ok(Some(line)) => line, Err(e) => return Err(e), }; match self.process_line(line)? { Some(item) => return Ok(Some(item)), None => continue, } } } /// Reads one complete line, without its ending. Returns `None` on a clean end with nothing /// buffered, and `Truncated` on a clean end with a partial line left over. fn read_line(&mut self) -> Result>, SseError> { loop { if let Some(pos) = self.buf.iter().position(|&b| b == b'\n') { let mut line: Vec = self.buf.drain(..=pos).collect(); line.pop(); // drop the '\n' if line.last() == Some(&b'\r') { line.pop(); } if line.len() > self.max_line { return Err(SseError::LineTooLong); } return Ok(Some(line)); } if self.buf.len() > self.max_line { return Err(SseError::LineTooLong); } let mut tmp = [0u8; 4096]; let n = match self.reader.read(&mut tmp) { Ok(0) => { if self.buf.is_empty() { return Ok(None); } return Err(SseError::Truncated); } Ok(n) => n, Err(e) if e.kind() == ErrorKind::Interrupted => continue, Err(e) if is_timeout(&e) => return Err(SseError::Timeout), Err(e) if e.kind() == ErrorKind::UnexpectedEof => return Err(SseError::Truncated), Err(e) => return Err(SseError::Io(e)), }; self.buf.extend_from_slice(&tmp[..n]); } } fn process_line(&self, line: Vec) -> Result, SseError> { let rest = match line.strip_prefix(b"data:") { Some(rest) => rest, None => return Ok(None), }; let rest = match rest.split_first() { Some((&b' ', rest)) => rest, _ => rest, }; let text = std::str::from_utf8(rest).map_err(|_| SseError::NotUtf8)?; if text == "[DONE]" { Ok(Some(SseItem::Done)) } else { Ok(Some(SseItem::Data(text.to_string()))) } } }