1use std::sync::{Arc, Mutex};
2
3use crate::types::{AssistantMessage, AssistantMessageEvent};
4
5#[derive(Clone)]
7pub struct AssistantMessageEventStream {
8 queue: Arc<Mutex<EventQueue>>,
9}
10
11const EVENT_COMPACT_THRESHOLD: usize = 64;
13
14struct EventQueue {
15 events: Vec<AssistantMessageEvent>,
16 read_index: usize,
17 done: bool,
18 final_result: Option<AssistantMessage>,
19 waiters: Vec<tokio::sync::oneshot::Sender<()>>,
20}
21
22fn compact_consumed_events(queue: &mut EventQueue) {
23 if queue.read_index >= EVENT_COMPACT_THRESHOLD {
24 queue.events.drain(0..queue.read_index);
25 queue.read_index = 0;
26 }
27}
28
29impl Default for AssistantMessageEventStream {
30 fn default() -> Self {
31 Self::new()
32 }
33}
34
35impl AssistantMessageEventStream {
36 pub fn new() -> Self {
37 Self {
38 queue: Arc::new(Mutex::new(EventQueue {
39 events: Vec::new(),
40 read_index: 0,
41 done: false,
42 final_result: None,
43 waiters: Vec::new(),
44 })),
45 }
46 }
47
48 pub fn clone_handle(&self) -> Self {
49 self.clone()
50 }
51
52 pub fn failed(message: impl Into<String>) -> Self {
53 let stream = Self::new();
54 let mut partial = AssistantMessage::empty(&crate::types::Model {
55 id: String::new(),
56 name: String::new(),
57 api: String::new(),
58 provider: String::new(),
59 base_url: String::new(),
60 reasoning: false,
61 thinking_level_map: None,
62 input: vec![],
63 cost: crate::types::ModelCost {
64 input: 0.0,
65 output: 0.0,
66 cache_read: 0.0,
67 cache_write: 0.0,
68
69 tiers: None,
70 },
71 context_window: 0,
72 max_tokens: 0,
73 headers: None,
74 openai_completions_compat: None,
75 openai_responses_compat: None,
76 anthropic_compat: None,
77 });
78 partial.stop_reason = crate::types::StopReason::Error;
79 partial.error_message = Some(message.into());
80 stream.push(AssistantMessageEvent::Error {
81 reason: crate::types::StopReason::Error,
82 error: partial,
83 });
84 stream.end();
85 stream
86 }
87
88 pub async fn next_event(&mut self) -> Option<AssistantMessageEvent> {
89 loop {
90 if let Some(event) = self.pop_next() {
91 return Some(event);
92 }
93 if self.is_done_sync() {
94 return None;
95 }
96 let rx = self.register_waiter();
97 let _ = rx.await;
98 }
99 }
100
101 pub fn is_done(&self) -> bool {
102 self.is_done_sync()
103 }
104
105 pub fn push(&self, event: AssistantMessageEvent) {
107 let mut q = self.queue.lock().expect("event stream mutex poisoned");
108 if q.done {
109 return;
110 }
111
112 match &event {
113 AssistantMessageEvent::Done { message, .. } => {
114 q.final_result = Some(message.clone());
115 q.done = true;
116 }
117 AssistantMessageEvent::Error { error, .. } => {
118 q.final_result = Some(error.clone());
119 q.done = true;
120 }
121 _ => {}
122 }
123
124 q.events.push(event);
125 let waiters = std::mem::take(&mut q.waiters);
126 for waiter in waiters {
127 let _ = waiter.send(());
128 }
129 }
130
131 pub fn end(&self) {
132 let mut q = self.queue.lock().expect("event stream mutex poisoned");
133 if q.done {
134 return;
135 }
136 q.done = true;
137 let waiters = std::mem::take(&mut q.waiters);
138 for waiter in waiters {
139 let _ = waiter.send(());
140 }
141 }
142
143 pub async fn result(&self) -> AssistantMessage {
144 loop {
145 if let Some(result) = self.final_result_sync() {
146 return result;
147 }
148 if self.is_done_sync() {
149 break;
150 }
151 let rx = self.register_waiter();
152 let _ = rx.await;
153 }
154 self.final_result_sync().unwrap_or_else(|| {
155 AssistantMessage::empty(&crate::types::Model {
156 id: String::new(),
157 name: String::new(),
158 api: String::new(),
159 provider: String::new(),
160 base_url: String::new(),
161 reasoning: false,
162 thinking_level_map: None,
163 input: vec![],
164 cost: crate::types::ModelCost {
165 input: 0.0,
166 output: 0.0,
167 cache_read: 0.0,
168 cache_write: 0.0,
169
170 tiers: None,
171 },
172 context_window: 0,
173 max_tokens: 0,
174 headers: None,
175 openai_completions_compat: None,
176 openai_responses_compat: None,
177 anthropic_compat: None,
178 })
179 })
180 }
181
182 fn pop_next(&self) -> Option<AssistantMessageEvent> {
183 let mut q = self.queue.lock().expect("event stream mutex poisoned");
184 if q.read_index < q.events.len() {
185 let event = q.events[q.read_index].clone();
186 q.read_index += 1;
187 compact_consumed_events(&mut q);
188 Some(event)
189 } else {
190 None
191 }
192 }
193
194 fn is_done_sync(&self) -> bool {
195 self.queue.lock().expect("event stream mutex poisoned").done
196 }
197
198 fn final_result_sync(&self) -> Option<AssistantMessage> {
199 self.queue
200 .lock()
201 .expect("event stream mutex poisoned")
202 .final_result
203 .clone()
204 }
205
206 fn register_waiter(&self) -> tokio::sync::oneshot::Receiver<()> {
207 let (tx, rx) = tokio::sync::oneshot::channel();
208 let mut q = self.queue.lock().expect("event stream mutex poisoned");
209 if q.read_index < q.events.len() || q.done {
210 let _ = tx.send(());
211 } else {
212 q.waiters.push(tx);
213 }
214 rx
215 }
216}
217
218pub struct EventStreamIterator {
219 queue: Arc<Mutex<EventQueue>>,
220 index: usize,
221}
222
223impl AssistantMessageEventStream {
224 pub fn into_stream(self) -> EventStreamIterator {
225 EventStreamIterator {
226 queue: self.queue,
227 index: 0,
228 }
229 }
230}
231
232impl EventStreamIterator {
233 pub async fn next(&mut self) -> Option<AssistantMessageEvent> {
234 loop {
235 let next = {
236 let mut q = self.queue.lock().expect("event stream mutex poisoned");
237 if self.index < q.events.len() {
238 let event = q.events[self.index].clone();
239 self.index += 1;
240 if self.index >= EVENT_COMPACT_THRESHOLD {
241 q.events.drain(0..self.index);
242 self.index = 0;
243 }
244 Some(event)
245 } else {
246 None
247 }
248 };
249 if next.is_some() {
250 return next;
251 }
252 let (done, register_waiter) = {
253 let q = self.queue.lock().expect("event stream mutex poisoned");
254 let done = q.done;
255 let register_waiter = self.index >= q.events.len() && !q.done;
256 (done, register_waiter)
257 };
258 if done {
259 return None;
260 }
261 if !register_waiter {
262 continue;
263 }
264 let rx = {
265 let (tx, rx) = tokio::sync::oneshot::channel();
266 let mut q = self.queue.lock().expect("event stream mutex poisoned");
267 if self.index < q.events.len() || q.done {
268 let _ = tx.send(());
269 } else {
270 q.waiters.push(tx);
271 }
272 rx
273 };
274 let _ = rx.await;
275 }
276 }
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282 use crate::types::{AssistantMessageEvent, Model, StopReason};
283
284 fn test_model() -> Model {
285 Model {
286 id: "test".to_string(),
287 name: "test".to_string(),
288 api: "test".to_string(),
289 provider: "test".to_string(),
290 base_url: "http://localhost".to_string(),
291 reasoning: false,
292 thinking_level_map: None,
293 input: vec!["text".to_string()],
294 cost: crate::types::ModelCost {
295 input: 0.0,
296 output: 0.0,
297 cache_read: 0.0,
298 cache_write: 0.0,
299
300 tiers: None,
301 },
302 context_window: 128_000,
303 max_tokens: 16_384,
304 headers: None,
305 openai_completions_compat: None,
306 openai_responses_compat: None,
307 anthropic_compat: None,
308 }
309 }
310
311 #[tokio::test]
312 async fn consumed_events_are_compacted_to_bound_memory() {
313 let stream = AssistantMessageEventStream::new();
314 let model = test_model();
315 for _ in 0..EVENT_COMPACT_THRESHOLD + 8 {
316 stream.push(AssistantMessageEvent::TextDelta {
317 content_index: 0,
318 delta: "x".to_string(),
319 partial: AssistantMessage::empty(&model),
320 });
321 }
322 stream.end();
323
324 let queue = stream.queue.clone();
325 let mut events = stream.into_stream();
326 let mut consumed = 0usize;
327 while events.next().await.is_some() {
328 consumed += 1;
329 }
330
331 let retained = queue.lock().expect("event stream mutex poisoned").events.len();
332 assert_eq!(consumed, EVENT_COMPACT_THRESHOLD + 8);
333 assert!(retained < EVENT_COMPACT_THRESHOLD);
334 }
335
336 #[tokio::test]
337 async fn waiter_is_not_registered_when_events_are_already_available() {
338 let mut stream = AssistantMessageEventStream::new();
339 let model = test_model();
340 let mut partial = AssistantMessage::empty(&model);
341 partial.stop_reason = StopReason::Stop;
342 stream.push(AssistantMessageEvent::Done {
343 reason: StopReason::Stop,
344 message: partial,
345 });
346 stream.end();
347
348 let event = stream.next_event().await.expect("stream event");
349 assert!(matches!(event, AssistantMessageEvent::Done { .. }));
350 }
351}