elph_ai/utils/
event_stream.rs1use std::sync::{Arc, Mutex};
2
3use crate::types::{AssistantMessage, AssistantMessageEvent};
4
5#[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 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}