Skip to main content

elph_ai/api/
sse.rs

1use anyhow::{Result, anyhow};
2use eventsource_stream::Eventsource;
3use futures::StreamExt;
4use serde_json::Value;
5use tokio_util::sync::CancellationToken;
6
7use super::common::request_aborted_error;
8
9/// Invoke `on_event` for each JSON event in an SSE response body.
10pub async fn for_each_sse_json_event<F>(
11    response: reqwest::Response,
12    signal: &Option<CancellationToken>,
13    mut on_event: F,
14) -> Result<()>
15where
16    F: FnMut(Value) -> Result<()>,
17{
18    let mut stream = response.bytes_stream().eventsource();
19    loop {
20        let next_item = match signal {
21            Some(token) => {
22                let token = token.clone();
23                tokio::select! {
24                    item = stream.next() => item,
25                    _ = token.cancelled() => return Err(request_aborted_error()),
26                }
27            }
28            None => stream.next().await,
29        };
30
31        let Some(item) = next_item else {
32            break;
33        };
34
35        let event = item?;
36        let data = event.data.trim();
37        if data.is_empty() || data == "[DONE]" {
38            continue;
39        }
40        let value =
41            serde_json::from_str::<Value>(data).map_err(|error| anyhow!("Invalid SSE JSON: {error}; data={data}"))?;
42        on_event(value)?;
43    }
44    Ok(())
45}
46
47/// Collect all JSON events from an SSE response body.
48pub async fn collect_sse_json_events(response: reqwest::Response) -> Result<Vec<Value>> {
49    let mut events = Vec::new();
50    for_each_sse_json_event(response, &None, |event| {
51        events.push(event);
52        Ok(())
53    })
54    .await?;
55    Ok(events)
56}
57
58/// Anthropic-style SSE decoder state.
59#[derive(Default)]
60pub struct SseDecoderState {
61    event: Option<String>,
62    data: Vec<String>,
63    raw: Vec<String>,
64}
65
66#[derive(Debug, Clone)]
67pub struct ServerSentEvent {
68    pub event: Option<String>,
69    pub data: String,
70    pub raw: Vec<String>,
71}
72
73pub(crate) fn flush_sse_event(state: &mut SseDecoderState) -> Option<ServerSentEvent> {
74    if state.event.is_none() && state.data.is_empty() {
75        return None;
76    }
77    let event = ServerSentEvent {
78        event: state.event.take(),
79        data: state.data.join("\n"),
80        raw: std::mem::take(&mut state.raw),
81    };
82    state.data.clear();
83    Some(event)
84}
85
86pub(crate) fn decode_sse_line(line: &str, state: &mut SseDecoderState) -> Option<ServerSentEvent> {
87    if line.is_empty() {
88        return flush_sse_event(state);
89    }
90    state.raw.push(line.to_string());
91    if line.starts_with(':') {
92        return None;
93    }
94    if let Some((field, value)) = line.split_once(':') {
95        let value = value.strip_prefix(' ').unwrap_or(value);
96        match field {
97            "event" => state.event = Some(value.to_string()),
98            "data" => state.data.push(value.to_string()),
99            _ => {}
100        }
101    }
102    None
103}
104
105/// Parse raw bytes into Anthropic SSE events.
106pub fn decode_sse_buffer(buffer: &str, state: &mut SseDecoderState) -> Vec<ServerSentEvent> {
107    let mut events = Vec::new();
108    for line in buffer.split('\n') {
109        let line = line.trim_end_matches('\r');
110        if let Some(event) = decode_sse_line(line, state) {
111            events.push(event);
112        }
113    }
114    events
115}
116
117fn process_sse_line_buffer(line_buffer: &mut String, state: &mut SseDecoderState) -> Vec<ServerSentEvent> {
118    let mut events = Vec::new();
119    while let Some(newline_pos) = line_buffer.find('\n') {
120        let line = line_buffer[..newline_pos].trim_end_matches('\r').to_string();
121        *line_buffer = line_buffer[newline_pos + 1..].to_string();
122        if let Some(event) = decode_sse_line(&line, state) {
123            events.push(event);
124        }
125    }
126    events
127}
128
129/// Invoke `on_event` for each Anthropic-style SSE event, reading the body incrementally.
130pub async fn for_each_anthropic_sse_event<F>(
131    response: reqwest::Response,
132    signal: &Option<CancellationToken>,
133    mut on_event: F,
134) -> Result<()>
135where
136    F: FnMut(ServerSentEvent) -> Result<()>,
137{
138    let mut state = SseDecoderState::default();
139    let mut line_buffer = String::new();
140    let mut byte_stream = response.bytes_stream();
141
142    loop {
143        let next_chunk = match signal {
144            Some(token) => {
145                let token = token.clone();
146                tokio::select! {
147                    item = byte_stream.next() => item,
148                    _ = token.cancelled() => return Err(request_aborted_error()),
149                }
150            }
151            None => byte_stream.next().await,
152        };
153
154        let Some(chunk) = next_chunk else {
155            break;
156        };
157
158        let chunk = chunk?;
159        line_buffer.push_str(&String::from_utf8_lossy(&chunk));
160        for event in process_sse_line_buffer(&mut line_buffer, &mut state) {
161            on_event(event)?;
162        }
163    }
164
165    if !line_buffer.is_empty()
166        && let Some(event) = decode_sse_line(line_buffer.trim_end_matches('\r'), &mut state)
167    {
168        on_event(event)?;
169    }
170    if let Some(event) = flush_sse_event(&mut state) {
171        on_event(event)?;
172    }
173    Ok(())
174}
175
176pub const ANTHROPIC_MESSAGE_EVENTS: &[&str] = &[
177    "message_start",
178    "message_delta",
179    "message_stop",
180    "content_block_start",
181    "content_block_delta",
182    "content_block_stop",
183];