use futures::StreamExt;
use crate::models::error::{BackendError, ModelError, Result};
use crate::models::stream::{StreamEvent, StreamSink, emit_all};
use crate::models::types::ModelResponse;
use crate::utils::{drain_complete_lines, drain_sse_events};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Framing {
Sse,
Ndjson,
}
impl Framing {
fn drain(self, buf: &mut Vec<u8>) -> Vec<String> {
match self {
Self::Sse => drain_sse_events(buf),
Self::Ndjson => drain_complete_lines(buf),
}
}
fn residue(self, buf: &[u8]) -> Option<String> {
match self {
Self::Sse => None,
Self::Ndjson => {
let tail = String::from_utf8_lossy(buf);
let trimmed = tail.trim();
(!trimmed.is_empty()).then(|| trimmed.to_string())
},
}
}
const fn cap_message(self) -> &'static str {
match self {
Self::Sse => "SSE stream exceeded {} byte reassembly cap without a complete event",
Self::Ndjson => "NDJSON stream exceeded {} byte reassembly cap without a complete line",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Flow {
Continue,
Stop,
}
pub trait StreamProtocol {
const FRAMING: Framing;
fn on_frame(&mut self, frame: &str, out: &mut Vec<StreamEvent>) -> Result<Flow>;
fn finish(self, out: &mut Vec<StreamEvent>) -> Result<ModelResponse>;
}
pub async fn drive_stream<P, S, B, E>(
mut body: S,
mut protocol: P,
sink: Option<&StreamSink>,
) -> Result<ModelResponse>
where
P: StreamProtocol,
S: futures::Stream<Item = std::result::Result<B, E>> + Unpin,
B: AsRef<[u8]>,
E: std::fmt::Display,
{
let mut buf: Vec<u8> = Vec::new();
let mut stopped = false;
'read: while let Some(chunk) = body.next().await {
let chunk = chunk.map_err(|e| ModelError::StreamError(e.to_string()))?;
if buf.len() > crate::constants::MAX_SSE_BUFFER_BYTES {
return Err(ModelError::StreamError(P::FRAMING.cap_message().replace(
"{}",
&crate::constants::MAX_SSE_BUFFER_BYTES.to_string(),
)));
}
buf.extend_from_slice(chunk.as_ref());
for frame in P::FRAMING.drain(&mut buf) {
if frame.trim().is_empty() {
continue;
}
let mut events: Vec<StreamEvent> = Vec::new();
let flow = protocol.on_frame(&frame, &mut events)?;
emit_all(sink, events).await?;
if flow == Flow::Stop {
stopped = true;
break 'read;
}
}
}
if !stopped && let Some(frame) = P::FRAMING.residue(&buf) {
let mut events: Vec<StreamEvent> = Vec::new();
protocol.on_frame(&frame, &mut events)?;
emit_all(sink, events).await?;
}
let mut events: Vec<StreamEvent> = Vec::new();
let response = protocol.finish(&mut events)?;
emit_all(sink, events).await?;
Ok(response)
}
pub async fn plain_http_error(response: reqwest::Response) -> ModelError {
let status = response.status().as_u16();
let debug = crate::models::error::ResponseDebugContext::from_headers(response.headers());
let message = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".to_string());
ModelError::Backend(BackendError::HttpError {
status,
message,
debug,
})
}
#[cfg(test)]
mod tests {
use super::*;
struct Recorder {
frames: Vec<String>,
stop_after: Option<usize>,
}
impl Recorder {
fn new() -> Self {
Self {
frames: Vec::new(),
stop_after: None,
}
}
}
struct SseRecorder(Recorder);
struct NdjsonRecorder(Recorder);
impl StreamProtocol for SseRecorder {
const FRAMING: Framing = Framing::Sse;
fn on_frame(&mut self, frame: &str, out: &mut Vec<StreamEvent>) -> Result<Flow> {
self.0.frames.push(frame.to_string());
out.push(StreamEvent::Text(frame.to_string()));
Ok(match self.0.stop_after {
Some(n) if self.0.frames.len() >= n => Flow::Stop,
_ => Flow::Continue,
})
}
fn finish(self, _out: &mut Vec<StreamEvent>) -> Result<ModelResponse> {
Ok(response_of(&self.0.frames))
}
}
impl StreamProtocol for NdjsonRecorder {
const FRAMING: Framing = Framing::Ndjson;
fn on_frame(&mut self, frame: &str, out: &mut Vec<StreamEvent>) -> Result<Flow> {
self.0.frames.push(frame.to_string());
out.push(StreamEvent::Text(frame.to_string()));
Ok(Flow::Continue)
}
fn finish(self, _out: &mut Vec<StreamEvent>) -> Result<ModelResponse> {
Ok(response_of(&self.0.frames))
}
}
fn response_of(frames: &[String]) -> ModelResponse {
ModelResponse {
content: frames.join("|"),
usage: None,
model_name: "recorder".to_string(),
stop_reason: None,
thinking: None,
tool_calls: None,
provider_continuation: None,
}
}
#[test]
fn sse_framing_drops_an_unterminated_tail() {
let mut buf = b"data: one\n\ndata: two".to_vec();
assert_eq!(Framing::Sse.drain(&mut buf), vec!["one".to_string()]);
assert_eq!(Framing::Sse.residue(&buf), None);
}
#[test]
fn ndjson_framing_keeps_an_unterminated_tail() {
let mut buf = b"{\"a\":1}\n{\"b\":2}".to_vec();
assert_eq!(
Framing::Ndjson.drain(&mut buf),
vec!["{\"a\":1}".to_string()]
);
assert_eq!(Framing::Ndjson.residue(&buf), Some("{\"b\":2}".to_string()));
}
#[test]
fn ndjson_residue_ignores_trailing_whitespace() {
assert_eq!(Framing::Ndjson.residue(b" \n "), None);
assert_eq!(Framing::Ndjson.residue(b""), None);
}
#[test]
fn cap_message_names_the_framing_it_protects() {
assert!(Framing::Sse.cap_message().starts_with("SSE"));
assert!(Framing::Ndjson.cap_message().starts_with("NDJSON"));
}
#[tokio::test]
async fn frames_split_across_chunks_reassemble() {
let (tx, mut rx) = tokio::sync::mpsc::channel::<StreamEvent>(16);
let done = tokio::spawn(async move {
drive_stream(
chunks(&["data: hel", "lo\n\ndata: wor", "ld\n\n"]),
SseRecorder(Recorder::new()),
Some(&tx),
)
.await
});
let mut seen = Vec::new();
while let Some(StreamEvent::Text(t)) = rx.recv().await {
seen.push(t);
}
let out = done.await.expect("join").expect("drive");
assert_eq!(seen, vec!["hello".to_string(), "world".to_string()]);
assert_eq!(out.content, "hello|world");
}
#[tokio::test]
async fn a_frame_delivered_one_byte_at_a_time_is_the_same_frame() {
let body = "data: hello\n\ndata: world\n\n";
let one_byte_each: Vec<Vec<u8>> = body.bytes().map(|b| vec![b]).collect();
let out = drive_stream(
byte_chunks(one_byte_each),
SseRecorder(Recorder::new()),
None,
)
.await
.expect("drive");
assert_eq!(out.content, "hello|world");
}
#[tokio::test]
async fn stop_ends_the_read_without_waiting_for_the_body() {
let mut recorder = Recorder::new();
recorder.stop_after = Some(1);
let out = drive_stream(
chunks(&["data: first\n\ndata: second\n\n"]),
SseRecorder(recorder),
None,
)
.await
.expect("drive");
assert_eq!(out.content, "first");
}
#[tokio::test]
async fn ndjson_flushes_the_frame_a_body_closed_on() {
let out = drive_stream(
chunks(&["{\"a\":1}\n{\"b\":2}"]),
NdjsonRecorder(Recorder::new()),
None,
)
.await
.expect("drive");
assert_eq!(out.content, "{\"a\":1}|{\"b\":2}");
}
#[tokio::test]
async fn sse_drops_the_frame_a_body_closed_on() {
let out = drive_stream(
chunks(&["data: whole\n\ndata: partial"]),
SseRecorder(Recorder::new()),
None,
)
.await
.expect("drive");
assert_eq!(out.content, "whole");
}
#[tokio::test]
async fn a_frameless_flood_trips_the_reassembly_cap() {
let filler = "x".repeat(crate::constants::MAX_SSE_BUFFER_BYTES / 2 + 16);
let err = drive_stream(
chunks(&[filler.as_str(), filler.as_str(), filler.as_str()]),
SseRecorder(Recorder::new()),
None,
)
.await
.expect_err("cap trips");
assert!(
matches!(&err, ModelError::StreamError(m) if m.contains("SSE stream exceeded")),
"expected an SSE cap StreamError, got {err:?}"
);
}
#[tokio::test]
async fn a_transport_failure_mid_body_is_a_stream_error() {
let body = futures::stream::iter(vec![
Ok::<Vec<u8>, String>(b"data: one\n\n".to_vec()),
Err("connection reset".to_string()),
]);
let err = drive_stream(body, SseRecorder(Recorder::new()), None)
.await
.expect_err("transport failure");
assert!(
matches!(&err, ModelError::StreamError(m) if m.contains("connection reset")),
"expected a transport StreamError, got {err:?}"
);
}
fn chunks(parts: &[&str]) -> impl futures::Stream<Item = std::result::Result<Vec<u8>, String>> {
byte_chunks(parts.iter().map(|p| p.as_bytes().to_vec()).collect())
}
fn byte_chunks(
parts: Vec<Vec<u8>>,
) -> impl futures::Stream<Item = std::result::Result<Vec<u8>, String>> {
futures::stream::iter(parts.into_iter().map(Ok))
}
}