Skip to main content

elph_ai/utils/
event_stream.rs

1use std::sync::{Arc, Mutex};
2
3use crate::types::{AssistantMessage, AssistantMessageEvent};
4
5/// Async event stream for assistant message streaming.
6#[derive(Clone)]
7pub struct AssistantMessageEventStream {
8    queue: Arc<Mutex<EventQueue>>,
9}
10
11/// Compact consumed prefix once this many events have been read.
12const 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    /// Push an event in-order. Must be synchronous to preserve stream ordering.
106    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}