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() {
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) -> Option<(&str, &str)> {
let idx = next_line_break_index(text)?;
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..]))
}
pub struct SseEventStream {
bytes_stream: futures::stream::BoxStream<'static, Result<Bytes, reqwest::Error>>,
state: SseDecoderState,
buffer: String,
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(),
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) {
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 {
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)) => {
return Err(AiError::Sse {
message: format!("error reading sse body: {e}"),
});
}
Some(Ok(chunk)) => {
self.buffer
.push_str(std::str::from_utf8(&chunk).unwrap_or(""));
}
}
}
}
}
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 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();
}
}