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
11struct EventQueue {
12    events: Vec<AssistantMessageEvent>,
13    read_index: usize,
14    done: bool,
15    final_result: Option<AssistantMessage>,
16    waiters: Vec<tokio::sync::oneshot::Sender<()>>,
17}
18
19impl Default for AssistantMessageEventStream {
20    fn default() -> Self {
21        Self::new()
22    }
23}
24
25impl AssistantMessageEventStream {
26    pub fn new() -> Self {
27        Self {
28            queue: Arc::new(Mutex::new(EventQueue {
29                events: Vec::new(),
30                read_index: 0,
31                done: false,
32                final_result: None,
33                waiters: Vec::new(),
34            })),
35        }
36    }
37
38    pub fn clone_handle(&self) -> Self {
39        self.clone()
40    }
41
42    pub fn failed(message: impl Into<String>) -> Self {
43        let stream = Self::new();
44        let mut partial = AssistantMessage::empty(&crate::types::Model {
45            id: String::new(),
46            name: String::new(),
47            api: String::new(),
48            provider: String::new(),
49            base_url: String::new(),
50            reasoning: false,
51            thinking_level_map: None,
52            input: vec![],
53            cost: crate::types::ModelCost {
54                input: 0.0,
55                output: 0.0,
56                cache_read: 0.0,
57                cache_write: 0.0,
58            },
59            context_window: 0,
60            max_tokens: 0,
61            headers: None,
62            openai_completions_compat: None,
63            openai_responses_compat: None,
64            anthropic_compat: None,
65        });
66        partial.stop_reason = crate::types::StopReason::Error;
67        partial.error_message = Some(message.into());
68        stream.push(AssistantMessageEvent::Error {
69            reason: crate::types::StopReason::Error,
70            error: partial,
71        });
72        stream.end();
73        stream
74    }
75
76    pub async fn next_event(&mut self) -> Option<AssistantMessageEvent> {
77        loop {
78            if let Some(event) = self.pop_next() {
79                return Some(event);
80            }
81            if self.is_done_sync() {
82                return None;
83            }
84            let rx = self.register_waiter();
85            let _ = rx.await;
86        }
87    }
88
89    pub fn is_done(&self) -> bool {
90        self.is_done_sync()
91    }
92
93    /// Push an event in-order. Must be synchronous to preserve stream ordering.
94    pub fn push(&self, event: AssistantMessageEvent) {
95        let mut q = self.queue.lock().expect("event stream mutex poisoned");
96        if q.done {
97            return;
98        }
99
100        match &event {
101            AssistantMessageEvent::Done { message, .. } => {
102                q.final_result = Some(message.clone());
103                q.done = true;
104            }
105            AssistantMessageEvent::Error { error, .. } => {
106                q.final_result = Some(error.clone());
107                q.done = true;
108            }
109            _ => {}
110        }
111
112        q.events.push(event);
113        let waiters = std::mem::take(&mut q.waiters);
114        for waiter in waiters {
115            let _ = waiter.send(());
116        }
117    }
118
119    pub fn end(&self) {
120        let mut q = self.queue.lock().expect("event stream mutex poisoned");
121        if q.done {
122            return;
123        }
124        q.done = true;
125        let waiters = std::mem::take(&mut q.waiters);
126        for waiter in waiters {
127            let _ = waiter.send(());
128        }
129    }
130
131    pub async fn result(&self) -> AssistantMessage {
132        loop {
133            if let Some(result) = self.final_result_sync() {
134                return result;
135            }
136            if self.is_done_sync() {
137                break;
138            }
139            let rx = self.register_waiter();
140            let _ = rx.await;
141        }
142        self.final_result_sync().unwrap_or_else(|| {
143            AssistantMessage::empty(&crate::types::Model {
144                id: String::new(),
145                name: String::new(),
146                api: String::new(),
147                provider: String::new(),
148                base_url: String::new(),
149                reasoning: false,
150                thinking_level_map: None,
151                input: vec![],
152                cost: crate::types::ModelCost {
153                    input: 0.0,
154                    output: 0.0,
155                    cache_read: 0.0,
156                    cache_write: 0.0,
157                },
158                context_window: 0,
159                max_tokens: 0,
160                headers: None,
161                openai_completions_compat: None,
162                openai_responses_compat: None,
163                anthropic_compat: None,
164            })
165        })
166    }
167
168    fn pop_next(&self) -> Option<AssistantMessageEvent> {
169        let mut q = self.queue.lock().expect("event stream mutex poisoned");
170        if q.read_index < q.events.len() {
171            let event = q.events[q.read_index].clone();
172            q.read_index += 1;
173            Some(event)
174        } else {
175            None
176        }
177    }
178
179    fn is_done_sync(&self) -> bool {
180        self.queue.lock().expect("event stream mutex poisoned").done
181    }
182
183    fn final_result_sync(&self) -> Option<AssistantMessage> {
184        self.queue
185            .lock()
186            .expect("event stream mutex poisoned")
187            .final_result
188            .clone()
189    }
190
191    fn register_waiter(&self) -> tokio::sync::oneshot::Receiver<()> {
192        let (tx, rx) = tokio::sync::oneshot::channel();
193        self.queue.lock().expect("event stream mutex poisoned").waiters.push(tx);
194        rx
195    }
196}
197
198pub struct EventStreamIterator {
199    queue: Arc<Mutex<EventQueue>>,
200    index: usize,
201}
202
203impl AssistantMessageEventStream {
204    pub fn into_stream(self) -> EventStreamIterator {
205        EventStreamIterator {
206            queue: self.queue,
207            index: 0,
208        }
209    }
210}
211
212impl EventStreamIterator {
213    pub async fn next(&mut self) -> Option<AssistantMessageEvent> {
214        loop {
215            let next = {
216                let q = self.queue.lock().expect("event stream mutex poisoned");
217                if self.index < q.events.len() {
218                    let event = q.events[self.index].clone();
219                    self.index += 1;
220                    Some(event)
221                } else {
222                    None
223                }
224            };
225            if next.is_some() {
226                return next;
227            }
228            if self.queue.lock().expect("event stream mutex poisoned").done {
229                return None;
230            }
231            let (tx, rx) = tokio::sync::oneshot::channel();
232            self.queue.lock().expect("event stream mutex poisoned").waiters.push(tx);
233            let _ = rx.await;
234        }
235    }
236}