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
9pub 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
47pub 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#[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
105pub 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
129pub 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];