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
11pub 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
49pub 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#[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
107pub 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
144pub 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];