Skip to main content

elph_ai/api/
sse.rs

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