use crate::error::AiError;
use crate::providers::anthropic::json_parse::parse_json_with_repair;
use bytes::Bytes;
use futures::StreamExt;
use tokio_util::sync::CancellationToken;
pub const ANTHROPIC_MESSAGE_EVENTS: &[&str] = &[
"message_start",
"message_delta",
"message_stop",
"content_block_start",
"content_block_delta",
"content_block_stop",
];
#[derive(Debug, Clone)]
pub struct ServerSentEvent {
pub event: Option<String>,
pub data: String,
pub raw: Vec<String>,
}
#[derive(Debug, Default)]
pub struct SseDecoderState {
event: Option<String>,
data: Vec<String>,
raw: Vec<String>,
}
impl SseDecoderState {
pub fn new() -> Self {
Self::default()
}
pub fn decode_line(&mut self, line: &str) -> Option<ServerSentEvent> {
if line.is_empty() {
return self.flush();
}
self.raw.push(line.to_string());
if line.starts_with(':') {
return None;
}
let (field, value) = match line.find(':') {
None => (line.to_string(), String::new()),
Some(idx) => {
let f = line[..idx].to_string();
let mut v = line[idx + 1..].to_string();
if let Some(stripped) = v.strip_prefix(' ') {
v = stripped.to_string();
}
(f, v)
}
};
if field == "event" {
self.event = Some(value);
} else if field == "data" {
self.data.push(value);
}
None
}
pub fn flush(&mut self) -> Option<ServerSentEvent> {
if self.event.is_none() && self.data.is_empty() {
self.raw.clear();
return None;
}
let event = ServerSentEvent {
event: self.event.take(),
data: self.data.join("\n"),
raw: std::mem::take(&mut self.raw),
};
self.data.clear();
Some(event)
}
}
fn next_line_break_index(text: &str) -> Option<usize> {
let cr = text.find('\r');
let nl = text.find('\n');
match (cr, nl) {
(None, None) => None,
(Some(a), None) => Some(a),
(None, Some(b)) => Some(b),
(Some(a), Some(b)) => Some(a.min(b)),
}
}
fn consume_line(text: &str, allow_trailing_cr: bool) -> Option<(&str, &str)> {
let idx = next_line_break_index(text)?;
if !allow_trailing_cr && idx + 1 == text.len() && text.as_bytes().get(idx) == Some(&b'\r') {
return None;
}
let mut next = idx + 1;
if text.as_bytes().get(idx) == Some(&b'\r') && text.as_bytes().get(next) == Some(&b'\n') {
next += 1;
}
Some((&text[..idx], &text[next..]))
}
fn append_utf8_chunk(buffer: &mut String, utf8_buffer: &mut Vec<u8>, chunk: &[u8]) {
utf8_buffer.extend_from_slice(chunk);
loop {
match std::str::from_utf8(utf8_buffer) {
Ok(text) => {
buffer.push_str(text);
utf8_buffer.clear();
break;
}
Err(error) => {
let valid = error.valid_up_to();
if valid > 0 {
let text = std::str::from_utf8(&utf8_buffer[..valid])
.expect("valid_up_to must identify valid UTF-8");
buffer.push_str(text);
utf8_buffer.drain(..valid);
continue;
}
if let Some(error_len) = error.error_len() {
buffer.push('\u{FFFD}');
utf8_buffer.drain(..error_len);
continue;
}
break;
}
}
}
}
pub struct SseEventStream {
bytes_stream: futures::stream::BoxStream<'static, Result<Bytes, reqwest::Error>>,
state: SseDecoderState,
buffer: String,
utf8_buffer: Vec<u8>,
signal: CancellationToken,
done: bool,
}
impl SseEventStream {
pub fn new(response: reqwest::Response, signal: CancellationToken) -> Self {
Self {
bytes_stream: response.bytes_stream().boxed(),
state: SseDecoderState::new(),
buffer: String::new(),
utf8_buffer: Vec::new(),
signal,
done: false,
}
}
pub async fn next_event(&mut self) -> Result<Option<ServerSentEvent>, AiError> {
loop {
let buffer = std::mem::take(&mut self.buffer);
let mut leftover = buffer;
let mut produced = None;
while let Some((line, rest)) = consume_line(&leftover, self.done) {
let line_owned = line.to_string();
leftover = rest.to_string();
if let Some(event) = self.state.decode_line(&line_owned) {
produced = Some(event);
break;
}
}
self.buffer = leftover;
if let Some(event) = produced {
return Ok(Some(event));
}
if self.done {
if !self.utf8_buffer.is_empty() {
self.buffer
.push_str(&String::from_utf8_lossy(&self.utf8_buffer));
self.utf8_buffer.clear();
}
let mut buffer = std::mem::take(&mut self.buffer);
if !buffer.is_empty() {
let pending = std::mem::take(&mut buffer);
if let Some(event) = self.state.decode_line(&pending) {
self.buffer = buffer;
return Ok(Some(event));
}
self.buffer = buffer;
}
return Ok(self.state.flush());
}
let next = tokio::select! {
biased;
_ = self.signal.cancelled() => {
return Err(AiError::Abort {
message: "Request was aborted".to_string(),
});
}
next = self.bytes_stream.next() => next,
};
match next {
None => {
self.done = true;
continue;
}
Some(Err(e)) => {
if self.signal.is_cancelled() {
return Err(AiError::Abort {
message: "Request was aborted".to_string(),
});
}
let mut source = String::new();
let mut current: &(dyn std::error::Error + 'static) = &e;
while let Some(next) = current.source() {
if !source.is_empty() {
source.push_str("; ");
}
source.push_str(&next.to_string());
current = next;
}
let detail = if source.is_empty() {
e.to_string()
} else {
format!("{} (caused by: {source})", e)
};
return Err(AiError::Sse {
message: format!("error reading sse body: {detail}"),
});
}
Some(Ok(chunk)) => {
append_utf8_chunk(&mut self.buffer, &mut self.utf8_buffer, &chunk);
}
}
}
}
}
pub async fn iterate_sse_messages(
response: reqwest::Response,
signal: &CancellationToken,
) -> Result<Vec<ServerSentEvent>, AiError> {
let mut stream = SseEventStream::new(response, signal.clone());
let mut out = Vec::new();
while let Some(event) = stream.next_event().await? {
out.push(event);
}
Ok(out)
}
pub fn parse_anthropic_event(sse: &ServerSentEvent) -> Result<AnthropicEvent, AiError> {
if sse.event.as_deref() == Some("error") {
return Err(AiError::Provider {
code: "sse_error".to_string(),
message: sse.data.clone(),
});
}
let event_name = match &sse.event {
Some(name) if ANTHROPIC_MESSAGE_EVENTS.contains(&name.as_str()) => name.clone(),
_ => return Ok(AnthropicEvent::Skipped),
};
let value: serde_json::Value = parse_json_with_repair(&sse.data).map_err(|e| AiError::Sse {
message: format!(
"Could not parse Anthropic SSE event {}: {}; data={}; raw={}",
event_name,
e,
sse.data,
sse.raw.join("\\n"),
),
})?;
Ok(AnthropicEvent::Message {
event_type: event_name,
payload: value,
})
}
#[derive(Debug, Clone)]
pub enum AnthropicEvent {
Message {
event_type: String,
payload: serde_json::Value,
},
Skipped,
}
impl AnthropicEvent {
pub fn event_type(&self) -> Option<&str> {
match self {
AnthropicEvent::Message { event_type, .. } => Some(event_type),
AnthropicEvent::Skipped => None,
}
}
pub fn payload(&self) -> Option<&serde_json::Value> {
match self {
AnthropicEvent::Message { payload, .. } => Some(payload),
AnthropicEvent::Skipped => None,
}
}
}
pub async fn iterate_anthropic_events(
response: reqwest::Response,
signal: &CancellationToken,
) -> Result<Vec<AnthropicEvent>, AiError> {
let frames = iterate_sse_messages(response, signal).await?;
let mut saw_message_start = false;
let mut saw_message_stop = false;
let mut out = Vec::with_capacity(frames.len());
for frame in frames {
let event = parse_anthropic_event(&frame)?;
match &event {
AnthropicEvent::Message { event_type, .. } => {
if event_type == "message_start" {
saw_message_start = true;
} else if event_type == "message_stop" {
saw_message_stop = true;
}
out.push(event);
}
AnthropicEvent::Skipped => {}
}
}
if saw_message_start && !saw_message_stop {
return Err(AiError::Sse {
message: "Anthropic stream ended before message_stop".to_string(),
});
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn decode_all(input: &str) -> Vec<ServerSentEvent> {
let mut state = SseDecoderState::new();
let mut out = Vec::new();
for line in input.split_inclusive('\n') {
let trimmed = line.trim_end_matches('\n').trim_end_matches('\r');
if let Some(event) = state.decode_line(trimmed) {
out.push(event);
}
}
if let Some(event) = state.flush() {
out.push(event);
}
out
}
#[test]
fn decodes_simple_message_start() {
let sse = "event: message_start\ndata: {\"type\":\"message_start\"}\n\n";
let events = decode_all(sse);
assert_eq!(events.len(), 1);
assert_eq!(events[0].event.as_deref(), Some("message_start"));
assert_eq!(events[0].data, "{\"type\":\"message_start\"}");
assert_eq!(events[0].raw.len(), 2);
}
#[test]
fn joins_multiline_data_with_newline() {
let sse = "event: content_block_delta\ndata: line1\ndata: line2\n\n";
let events = decode_all(sse);
assert_eq!(events[0].data, "line1\nline2");
}
#[test]
fn comment_lines_are_dropped() {
let sse = ": keepalive\nevent: ping\ndata: {}\n\nevent: message_start\ndata: {}\n\n";
let events = decode_all(sse);
assert_eq!(events.len(), 2);
assert_eq!(events[0].event.as_deref(), Some("ping"));
assert_eq!(events[1].event.as_deref(), Some("message_start"));
}
#[test]
fn strips_single_leading_space_after_colon() {
let sse = "event: message_delta\ndata: {\"x\":1}\n\n";
let events = decode_all(sse);
assert_eq!(events[0].data, "{\"x\":1}");
}
#[test]
fn keeps_trailing_cr_until_the_next_chunk() {
assert!(consume_line("data: value\r", false).is_none());
let (line, rest) = consume_line("data: value\r\n", false).unwrap();
assert_eq!(line, "data: value");
assert!(rest.is_empty());
}
#[test]
fn eof_consumes_a_standalone_cr_and_flushes_the_last_event() {
let mut state = SseDecoderState::new();
let mut buffer = "event: message_start\rdata: {}\r".to_string();
let mut events = Vec::new();
while let Some((line, rest)) = consume_line(&buffer, true) {
let line = line.to_string();
buffer = rest.to_string();
if let Some(event) = state.decode_line(&line) {
events.push(event);
}
}
assert!(buffer.is_empty());
if let Some(event) = state.flush() {
events.push(event);
}
assert_eq!(events.len(), 1);
assert_eq!(events[0].event.as_deref(), Some("message_start"));
assert_eq!(events[0].data, "{}");
}
#[test]
fn split_utf8_code_point_is_reassembled_across_chunks() {
let mut buffer = String::new();
let mut pending = Vec::new();
append_utf8_chunk(&mut buffer, &mut pending, b"data: \xe4");
assert_eq!(buffer, "data: ");
assert_eq!(pending, vec![0xe4]);
append_utf8_chunk(&mut buffer, &mut pending, b"\xb8\xad\n\n");
assert_eq!(buffer, "data: \u{4e2d}\n\n");
assert!(pending.is_empty());
}
#[test]
fn malformed_utf8_does_not_block_following_sse_bytes() {
let mut buffer = String::new();
let mut pending = Vec::new();
append_utf8_chunk(&mut buffer, &mut pending, b"data: \xff");
append_utf8_chunk(&mut buffer, &mut pending, b"ok\n\n");
assert_eq!(buffer, "data: \u{fffd}ok\n\n");
assert!(pending.is_empty());
}
#[test]
fn nameless_event_has_none_event_field() {
let sse = "data: only-data\n\n";
let events = decode_all(sse);
assert_eq!(events.len(), 1);
assert!(events[0].event.is_none());
}
#[test]
fn error_event_surfaces_provider_error() {
let frame = ServerSentEvent {
event: Some("error".to_string()),
data: "rate limited".to_string(),
raw: vec!["data: rate limited".to_string()],
};
let err = parse_anthropic_event(&frame).unwrap_err();
assert!(matches!(err, AiError::Provider { code, .. } if code == "sse_error"));
}
#[test]
fn non_message_event_is_skipped() {
let frame = ServerSentEvent {
event: Some("ping".to_string()),
data: "{}".to_string(),
raw: vec!["data: {}".to_string()],
};
let event = parse_anthropic_event(&frame).unwrap();
assert!(matches!(event, AnthropicEvent::Skipped));
}
#[test]
fn malformed_message_data_is_sse_error() {
let frame = ServerSentEvent {
event: Some("message_start".to_string()),
data: "not json".to_string(),
raw: vec!["data: not json".to_string()],
};
let err = parse_anthropic_event(&frame).unwrap_err();
assert!(matches!(err, AiError::Sse { .. }));
}
#[tokio::test]
async fn cancellation_interrupts_pending_body_read() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 1024];
let _ = socket.read(&mut request).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Transfer-Encoding: chunked\r\nConnection: keep-alive\r\n\r\n",
)
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
});
let response = reqwest::Client::new()
.get(format!("http://{address}"))
.send()
.await
.unwrap();
let signal = CancellationToken::new();
let cancel = signal.clone();
let mut stream = SseEventStream::new(response, signal);
tokio::spawn(async move {
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
cancel.cancel();
});
let result = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next_event())
.await
.expect("cancellation must wake a pending response body read");
assert!(matches!(result, Err(AiError::Abort { .. })));
server.abort();
}
#[tokio::test]
async fn body_decode_error_keeps_underlying_transport_detail() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 1024];
let _ = socket.read(&mut request).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Transfer-Encoding: chunked\r\nConnection: close\r\n\r\nZZ\r\n",
)
.await
.unwrap();
});
let response = reqwest::Client::new()
.get(format!("http://{address}"))
.send()
.await
.unwrap();
let mut stream = SseEventStream::new(response, CancellationToken::new());
let error = stream.next_event().await.unwrap_err();
let message = error.to_string();
assert!(message.contains("error reading sse body"));
assert!(message.contains("error decoding response body"));
assert!(message.contains("missing size digit"), "{message}");
server.await.unwrap();
}
#[tokio::test]
async fn body_timeout_is_reported_with_timeout_detail() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0u8; 1024];
let _ = socket.read(&mut request).await;
socket
.write_all(
b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\n\
Transfer-Encoding: chunked\r\nConnection: keep-alive\r\n\r\n",
)
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
});
let response = reqwest::Client::builder()
.timeout(std::time::Duration::from_millis(30))
.build()
.unwrap()
.get(format!("http://{address}"))
.send()
.await
.unwrap();
let mut stream = SseEventStream::new(response, CancellationToken::new());
let error = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next_event())
.await
.unwrap()
.unwrap_err();
let message = error.to_string();
assert!(
message.contains("error decoding response body"),
"{message}"
);
assert!(message.to_lowercase().contains("timed out"), "{message}");
server.abort();
}
}