gemini_adk_rs/live/
effect_executor.rs

1//! Execute typed Live reactor effects against a session writer.
2
3use 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/// Executes [`LiveEffect`] values emitted by the Live reactor.
15#[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    /// Create an executor backed by a session writer.
24    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    /// Execute a list of policy-wrapped reactions.
37    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                        // Supervise: surface concurrent failures rather than
72                        // silently dropping them.
73                        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    /// Execute one typed effect.
86    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    /// Flush deferred context and an armed prompt.
115    ///
116    /// This is intentionally gated by [`PendingContext::take_prompt`], so a
117    /// playback-drained event cannot trigger a new empty model turn unless the
118    /// control plane explicitly armed one.
119    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        // A concurrent effect that fails must surface as a LiveEvent, not vanish.
354        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}