1use std::sync::Arc;
4
5use gemini_genai_rs::prelude::SessionPhase;
6use gemini_genai_rs::session::{SessionError, SessionWriter};
7use tokio::sync::broadcast;
8
9use super::ExecutionMode;
10use super::context_writer::PendingContext;
11use super::events::LiveEvent;
12use super::reactor::{LiveEffect, Reaction};
13
14#[derive(Clone)]
16pub struct LiveEffectExecutor {
17 writer: Arc<dyn SessionWriter>,
18 pending_context: Option<Arc<PendingContext>>,
19 event_tx: broadcast::Sender<LiveEvent>,
20}
21
22impl LiveEffectExecutor {
23 pub fn new(
25 writer: Arc<dyn SessionWriter>,
26 pending_context: Option<Arc<PendingContext>>,
27 event_tx: broadcast::Sender<LiveEvent>,
28 ) -> Self {
29 Self {
30 writer,
31 pending_context,
32 event_tx,
33 }
34 }
35
36 pub async fn execute_reactions(&self, reactions: Vec<Reaction>) -> Result<(), SessionError> {
38 for reaction in reactions {
39 match reaction.policy.mode {
40 ExecutionMode::Blocking => {
41 let executor = self.clone();
42 let fut = executor.execute(reaction.effect);
43 if let Some(timeout) = reaction.policy.timeout {
44 tokio::time::timeout(timeout, fut).await.map_err(|_| {
45 SessionError::Timeout {
46 phase: SessionPhase::Active,
47 elapsed: timeout,
48 }
49 })??;
50 } else {
51 fut.await?;
52 }
53 }
54 ExecutionMode::Concurrent => {
55 let executor = self.clone();
56 let timeout = reaction.policy.timeout;
57 let source = reaction.source;
58 let effect = reaction.effect;
59 tokio::spawn(async move {
60 let result = match timeout {
61 Some(timeout) => {
62 tokio::time::timeout(timeout, executor.execute(effect))
63 .await
64 .unwrap_or(Err(SessionError::Timeout {
65 phase: SessionPhase::Active,
66 elapsed: timeout,
67 }))
68 }
69 None => executor.execute(effect).await,
70 };
71 if let Err(err) = result {
74 let _ = executor.event_tx.send(LiveEvent::Error(format!(
75 "reaction '{source}' failed: {err}"
76 )));
77 }
78 });
79 }
80 }
81 }
82 Ok(())
83 }
84
85 pub async fn execute(&self, effect: LiveEffect) -> Result<(), SessionError> {
87 match effect {
88 LiveEffect::Noop => Ok(()),
89 LiveEffect::SendContext(contents) => {
90 if !contents.is_empty() {
91 self.writer.send_client_content(contents, false).await?;
92 }
93 Ok(())
94 }
95 LiveEffect::PromptModel => self.flush_deferred_prompt().await,
96 LiveEffect::CancelDeferredPrompt => {
97 if let Some(pending) = &self.pending_context {
98 pending.clear_prompt();
99 }
100 Ok(())
101 }
102 LiveEffect::SignalUserActivityStart => self.writer.signal_activity_start().await,
103 LiveEffect::SignalUserActivityEnd => self.writer.signal_activity_end().await,
104 LiveEffect::UpdateInstruction(instruction) => {
105 self.writer.update_instruction(instruction).await
106 }
107 LiveEffect::Emit(event) => {
108 let _ = self.event_tx.send(event);
109 Ok(())
110 }
111 }
112 }
113
114 pub async fn flush_deferred_prompt(&self) -> Result<(), SessionError> {
120 let Some(pending) = &self.pending_context else {
121 return Ok(());
122 };
123
124 let contents = pending.drain_context();
125 if !contents.is_empty() {
126 self.writer.send_client_content(contents, false).await?;
127 }
128 if pending.take_prompt() {
129 self.writer.send_client_content(vec![], true).await?;
130 }
131 Ok(())
132 }
133}
134
135#[cfg(test)]
136mod tests {
137 use super::*;
138 use async_trait::async_trait;
139 use gemini_genai_rs::prelude::{Content, FunctionResponse};
140 use parking_lot::Mutex;
141
142 #[derive(Debug, Clone, PartialEq, Eq)]
143 enum Write {
144 ClientContent { turns: usize, turn_complete: bool },
145 Instruction(String),
146 ActivityStart,
147 ActivityEnd,
148 }
149
150 #[derive(Default)]
151 struct MockWriter {
152 writes: Mutex<Vec<Write>>,
153 }
154
155 #[async_trait]
156 impl SessionWriter for MockWriter {
157 async fn send_audio(&self, _data: bytes::Bytes) -> Result<(), SessionError> {
158 Ok(())
159 }
160
161 async fn send_text(&self, _text: String) -> Result<(), SessionError> {
162 Ok(())
163 }
164
165 async fn send_tool_response(
166 &self,
167 _responses: Vec<FunctionResponse>,
168 ) -> Result<(), SessionError> {
169 Ok(())
170 }
171
172 async fn send_client_content(
173 &self,
174 turns: Vec<Content>,
175 turn_complete: bool,
176 ) -> Result<(), SessionError> {
177 self.writes.lock().push(Write::ClientContent {
178 turns: turns.len(),
179 turn_complete,
180 });
181 Ok(())
182 }
183
184 async fn send_video(&self, _jpeg_data: bytes::Bytes) -> Result<(), SessionError> {
185 Ok(())
186 }
187
188 async fn update_instruction(&self, instruction: String) -> Result<(), SessionError> {
189 self.writes.lock().push(Write::Instruction(instruction));
190 Ok(())
191 }
192
193 async fn signal_activity_start(&self) -> Result<(), SessionError> {
194 self.writes.lock().push(Write::ActivityStart);
195 Ok(())
196 }
197
198 async fn signal_activity_end(&self) -> Result<(), SessionError> {
199 self.writes.lock().push(Write::ActivityEnd);
200 Ok(())
201 }
202
203 async fn disconnect(&self) -> Result<(), SessionError> {
204 Ok(())
205 }
206 }
207
208 #[tokio::test]
209 async fn prompt_model_flushes_context_then_armed_prompt() {
210 let writer = Arc::new(MockWriter::default());
211 let pending = Arc::new(PendingContext::new());
212 pending.push(Content::model("phase context"));
213 pending.set_prompt();
214 let (event_tx, _) = broadcast::channel(8);
215 let executor = LiveEffectExecutor::new(writer.clone(), Some(pending.clone()), event_tx);
216
217 executor.execute(LiveEffect::PromptModel).await.unwrap();
218
219 assert_eq!(
220 writer.writes.lock().as_slice(),
221 &[
222 Write::ClientContent {
223 turns: 1,
224 turn_complete: false
225 },
226 Write::ClientContent {
227 turns: 0,
228 turn_complete: true
229 }
230 ]
231 );
232 assert!(pending.is_empty());
233 }
234
235 #[tokio::test]
236 async fn prompt_model_without_armed_prompt_only_flushes_context() {
237 let writer = Arc::new(MockWriter::default());
238 let pending = Arc::new(PendingContext::new());
239 pending.push(Content::model("phase context"));
240 let (event_tx, _) = broadcast::channel(8);
241 let executor = LiveEffectExecutor::new(writer.clone(), Some(pending), event_tx);
242
243 executor.execute(LiveEffect::PromptModel).await.unwrap();
244
245 assert_eq!(
246 writer.writes.lock().as_slice(),
247 &[Write::ClientContent {
248 turns: 1,
249 turn_complete: false
250 }]
251 );
252 }
253
254 #[tokio::test]
255 async fn update_instruction_uses_writer() {
256 let writer = Arc::new(MockWriter::default());
257 let (event_tx, _) = broadcast::channel(8);
258 let executor = LiveEffectExecutor::new(writer.clone(), None, event_tx);
259
260 executor
261 .execute(LiveEffect::UpdateInstruction("new instruction".into()))
262 .await
263 .unwrap();
264
265 assert_eq!(
266 writer.writes.lock().as_slice(),
267 &[Write::Instruction("new instruction".into())]
268 );
269 }
270
271 #[tokio::test]
272 async fn cancel_deferred_prompt_keeps_context() {
273 let writer = Arc::new(MockWriter::default());
274 let pending = Arc::new(PendingContext::new());
275 pending.push(Content::model("still useful with user audio"));
276 pending.set_prompt();
277 let (event_tx, _) = broadcast::channel(8);
278 let executor = LiveEffectExecutor::new(writer, Some(pending.clone()), event_tx);
279
280 executor
281 .execute(LiveEffect::CancelDeferredPrompt)
282 .await
283 .unwrap();
284
285 assert!(!pending.has_prompt());
286 assert_eq!(pending.drain_context().len(), 1);
287 }
288
289 #[tokio::test]
290 async fn user_activity_effects_signal_writer() {
291 let writer = Arc::new(MockWriter::default());
292 let (event_tx, _) = broadcast::channel(8);
293 let executor = LiveEffectExecutor::new(writer.clone(), None, event_tx);
294
295 executor
296 .execute_reactions(vec![
297 Reaction::blocking("test", LiveEffect::SignalUserActivityStart),
298 Reaction::blocking("test", LiveEffect::SignalUserActivityEnd),
299 ])
300 .await
301 .unwrap();
302
303 assert_eq!(
304 writer.writes.lock().as_slice(),
305 &[Write::ActivityStart, Write::ActivityEnd]
306 );
307 }
308
309 #[tokio::test]
310 async fn concurrent_effect_failure_is_surfaced_as_event() {
311 struct FailWriter;
312 #[async_trait]
313 impl SessionWriter for FailWriter {
314 async fn send_audio(&self, _: bytes::Bytes) -> Result<(), SessionError> {
315 Ok(())
316 }
317 async fn send_text(&self, _: String) -> Result<(), SessionError> {
318 Ok(())
319 }
320 async fn send_tool_response(
321 &self,
322 _: Vec<FunctionResponse>,
323 ) -> Result<(), SessionError> {
324 Ok(())
325 }
326 async fn send_client_content(
327 &self,
328 _: Vec<Content>,
329 _: bool,
330 ) -> Result<(), SessionError> {
331 Err(SessionError::NotConnected)
332 }
333 async fn send_video(&self, _: bytes::Bytes) -> Result<(), SessionError> {
334 Ok(())
335 }
336 async fn update_instruction(&self, _: String) -> Result<(), SessionError> {
337 Ok(())
338 }
339 async fn signal_activity_start(&self) -> Result<(), SessionError> {
340 Ok(())
341 }
342 async fn signal_activity_end(&self) -> Result<(), SessionError> {
343 Ok(())
344 }
345 async fn disconnect(&self) -> Result<(), SessionError> {
346 Ok(())
347 }
348 }
349
350 let (event_tx, mut rx) = broadcast::channel(8);
351 let executor = LiveEffectExecutor::new(Arc::new(FailWriter), None, event_tx);
352
353 executor
355 .execute_reactions(vec![Reaction::concurrent(
356 "test",
357 LiveEffect::SendContext(vec![Content::model("x")]),
358 )])
359 .await
360 .unwrap();
361
362 let event = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv())
363 .await
364 .expect("a reaction-failure event within the timeout")
365 .expect("event received");
366 assert!(
367 matches!(&event, LiveEvent::Error(msg) if msg.contains("reaction 'test' failed")),
368 "expected a reaction-failure error event, got {event:?}"
369 );
370 }
371}