gemini_adk_rs/live/
session_signals.rs

1//! Auto-tracked session-level state signals.
2//!
3//! [`SessionSignals`] is called by the telemetry lane on every
4//! [`SessionEvent`] and transparently updates keys under the `session:`
5//! prefix in the shared [`State`].
6//!
7//! Hot-path timestamps use [`AtomicU64`] (nanos since start) instead of
8//! `Mutex<Instant>`, eliminating per-event mutex contention. Derived
9//! timing signals (`silence_ms`, `elapsed_ms`, `remaining_budget_ms`)
10//! are flushed periodically via `flush_timing()` rather than on every event.
11
12use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
13use std::time::Instant;
14
15use gemini_genai_rs::prelude::{SessionEvent, SessionPhase};
16use parking_lot::Mutex;
17
18use crate::state::State;
19
20// ---------------------------------------------------------------------------
21// Public types
22// ---------------------------------------------------------------------------
23
24/// Session type determines the server-side duration limit.
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum SessionType {
27    /// Audio-only session (~15 min limit).
28    AudioOnly,
29    /// Audio + video session (~2 min limit).
30    AudioVideo,
31}
32
33// ---------------------------------------------------------------------------
34// SessionSignals
35// ---------------------------------------------------------------------------
36
37/// Tracks session-level signals automatically from events.
38///
39/// Every call to [`on_event`](SessionSignals::on_event) updates the
40/// corresponding keys under `session:` in the shared [`State`], making
41/// them available to instruction templates, watchers, and computed vars.
42///
43/// **Performance**: Timestamps use `AtomicU64` (nanos since session start)
44/// instead of `Mutex<Instant>`. Derived timing signals are flushed
45/// periodically via `flush_timing()` (100ms interval) rather than per-event.
46pub struct SessionSignals {
47    state: State,
48    /// Session start time — used as epoch for all atomic timestamps.
49    start: Instant,
50    /// Nanos since start when connected.
51    connected_at_ns: AtomicU64,
52    /// Whether currently connected.
53    is_connected: AtomicBool,
54    /// Nanos since start of last activity.
55    last_activity_ns: AtomicU64,
56    /// Whether the session includes video input.
57    has_video: AtomicBool,
58    /// Server-sent GoAway timestamp, if received.
59    go_away_at: Mutex<Option<Instant>>,
60    /// Latest resumption handle from server (persisted for reconnection).
61    latest_resume_handle: Mutex<Option<String>>,
62}
63
64impl SessionSignals {
65    /// Create a new `SessionSignals` backed by the given [`State`].
66    pub fn new(state: State) -> Self {
67        Self {
68            state,
69            start: Instant::now(),
70            connected_at_ns: AtomicU64::new(0),
71            is_connected: AtomicBool::new(false),
72            last_activity_ns: AtomicU64::new(0),
73            has_video: AtomicBool::new(false),
74            go_away_at: Mutex::new(None),
75            latest_resume_handle: Mutex::new(None),
76        }
77    }
78
79    /// Process an event — updates state keys and atomic timestamps.
80    ///
81    /// This is the per-event handler. It updates boolean flags, counters,
82    /// and atomic timestamps. **Derived timing** (silence_ms, elapsed_ms,
83    /// remaining_budget_ms) is NOT computed here — call `flush_timing()`
84    /// periodically instead.
85    pub fn on_event(&self, event: &SessionEvent) {
86        match event {
87            SessionEvent::Connected => {
88                let now_ns = self.elapsed_ns();
89                self.connected_at_ns.store(now_ns, Ordering::Relaxed);
90                self.is_connected.store(true, Ordering::Relaxed);
91                self.last_activity_ns.store(now_ns, Ordering::Relaxed);
92                let _ = self.state.session().set("connected_at_ms", 0u64);
93                let _ = self.state.session().set("interrupt_count", 0u64);
94                let _ = self.state.session().set("error_count", 0u64);
95                let _ = self.state.session().set("is_user_speaking", false);
96                let _ = self.state.session().set("is_model_speaking", false);
97                let _ = self.state.session().set("go_away_received", false);
98                let _ = self.state.session().set("resumable", false);
99                let _ = self.state.session().set("session_type", "audio_only");
100            }
101
102            SessionEvent::VoiceActivityStart => {
103                let _ = self.state.session().set("is_user_speaking", true);
104                self.touch_activity();
105            }
106
107            SessionEvent::VoiceActivityEnd => {
108                let _ = self.state.session().set("is_user_speaking", false);
109                self.touch_activity();
110            }
111
112            SessionEvent::Interrupted => {
113                let count: u64 = self.state.session().get("interrupt_count").unwrap_or(0);
114                let _ = self.state.session().set("interrupt_count", count + 1);
115                self.touch_activity();
116            }
117
118            SessionEvent::Error(msg) => {
119                let count: u64 = self.state.session().get("error_count").unwrap_or(0);
120                let _ = self.state.session().set("error_count", count + 1);
121                let _ = self.state.session().set("last_error", msg.to_string());
122            }
123
124            SessionEvent::PhaseChanged(phase) => {
125                let _ = self
126                    .state
127                    .session()
128                    .set("is_model_speaking", *phase == SessionPhase::ModelSpeaking);
129                let _ = self.state.session().set("phase", phase.to_string());
130                self.touch_activity();
131            }
132
133            SessionEvent::GoAway(time_left) => {
134                let _ = self.state.session().set("go_away_received", true);
135                if let Some(tl) = time_left {
136                    *self.go_away_at.lock() = Some(Instant::now() + *tl);
137                    let _ = self
138                        .state
139                        .session()
140                        .set("go_away_time_left_ms", tl.as_millis() as u64);
141                }
142            }
143
144            SessionEvent::SessionResumeUpdate(info) => {
145                *self.latest_resume_handle.lock() = Some(info.handle.clone());
146                let _ = self.state.session().set("resumable", info.resumable);
147                if let Some(ref idx) = info.last_consumed_index {
148                    let _ = self
149                        .state
150                        .session()
151                        .set("last_consumed_client_index", idx.clone());
152                }
153            }
154
155            SessionEvent::Usage(usage) => {
156                if let Some(total) = usage.total_token_count {
157                    let _ = self.state.session().set("total_token_count", total);
158                }
159                if let Some(prompt) = usage.prompt_token_count {
160                    let _ = self.state.session().set("prompt_token_count", prompt);
161                }
162                if let Some(response) = usage.response_token_count {
163                    let _ = self.state.session().set("response_token_count", response);
164                }
165                if let Some(cached) = usage.cached_content_token_count {
166                    let _ = self
167                        .state
168                        .session()
169                        .set("cached_content_token_count", cached);
170                }
171                if let Some(thoughts) = usage.thoughts_token_count {
172                    let _ = self.state.session().set("thoughts_token_count", thoughts);
173                }
174            }
175
176            SessionEvent::GenerationComplete => {
177                // No-op for signals — generation complete is handled by control lane
178            }
179
180            SessionEvent::InputTranscription(text) => {
181                let _ = self
182                    .state
183                    .session()
184                    .set("last_input_transcription", text.clone());
185                self.touch_activity();
186            }
187
188            SessionEvent::OutputTranscription(text) => {
189                let _ = self
190                    .state
191                    .session()
192                    .set("last_output_transcription", text.clone());
193                self.touch_activity();
194            }
195
196            SessionEvent::AudioData(_)
197            | SessionEvent::TextDelta(_)
198            | SessionEvent::TextComplete(_) => {
199                // High-frequency events: only touch the atomic timestamp.
200                // No DashMap writes, no mutex locks.
201                self.touch_activity();
202            }
203
204            SessionEvent::TurnComplete => {
205                self.touch_activity();
206            }
207
208            SessionEvent::Disconnected(_reason) => {
209                self.is_connected.store(false, Ordering::Relaxed);
210                let _ = self.state.session().set("disconnected", true);
211            }
212
213            _ => {}
214        }
215    }
216
217    /// Flush derived timing signals to state.
218    ///
219    /// Call this periodically (e.g., every 100ms) from the telemetry lane.
220    /// Computes `silence_ms`, `elapsed_ms`, and `remaining_budget_ms` from
221    /// atomic timestamps without any mutex locks.
222    pub fn flush_timing(&self) {
223        let last_activity = self.last_activity_ns.load(Ordering::Relaxed);
224        if last_activity > 0 {
225            let now_ns = self.elapsed_ns();
226            let silence_ms = now_ns.saturating_sub(last_activity) / 1_000_000;
227            let _ = self.state.session().set("silence_ms", silence_ms);
228        }
229
230        if self.is_connected.load(Ordering::Relaxed) {
231            let connected_ns = self.connected_at_ns.load(Ordering::Relaxed);
232            let now_ns = self.elapsed_ns();
233            let elapsed_ms = now_ns.saturating_sub(connected_ns) / 1_000_000;
234            let _ = self.state.session().set("elapsed_ms", elapsed_ms);
235
236            let limit_ms: u64 = match self.session_type() {
237                SessionType::AudioOnly => 15 * 60 * 1000,
238                SessionType::AudioVideo => 2 * 60 * 1000,
239            };
240            let remaining = limit_ms.saturating_sub(elapsed_ms);
241            let _ = self.state.session().set("remaining_budget_ms", remaining);
242        }
243    }
244
245    #[inline]
246    fn touch_activity(&self) {
247        self.last_activity_ns
248            .store(self.elapsed_ns(), Ordering::Relaxed);
249    }
250
251    #[inline]
252    fn elapsed_ns(&self) -> u64 {
253        self.start.elapsed().as_nanos() as u64
254    }
255
256    /// Record the turn's response latency (end of user speech, or text send,
257    /// to the model's first output) as `session:last_response_latency_ms`, so
258    /// instruction templates, watchers and computed vars can react to a slow
259    /// turn. Called by the telemetry lane once per measured turn.
260    ///
261    /// The write happens when the turn's *first* output event is observed, so
262    /// for a spoken turn it lands a whole model response before the boundary.
263    /// The telemetry and control lanes are independent readers of the same
264    /// event stream, though, with no barrier between them: if the telemetry
265    /// lane is running behind, a turn-boundary reader can still see the
266    /// previous turn's value. Treat the key as the latest measurement, not as
267    /// one synchronised to the current turn.
268    pub fn record_response_latency(&self, latency: std::time::Duration) {
269        let _ = self
270            .state
271            .session()
272            .set("last_response_latency_ms", latency.as_millis() as u64);
273        self.touch_activity();
274    }
275
276    /// Returns the current session type based on whether video has been sent.
277    pub fn session_type(&self) -> SessionType {
278        if self.has_video.load(Ordering::Relaxed) {
279            SessionType::AudioVideo
280        } else {
281            SessionType::AudioOnly
282        }
283    }
284
285    /// Returns the latest resumption handle for reconnection.
286    pub fn latest_resume_handle(&self) -> Option<String> {
287        self.latest_resume_handle.lock().clone()
288    }
289
290    /// Mark that video has been sent (changes session type to `AudioVideo`).
291    pub fn mark_video_sent(&self) {
292        if !self.has_video.swap(true, Ordering::Relaxed) {
293            let _ = self.state.session().set("session_type", "audio_video");
294        }
295    }
296}
297
298#[cfg(test)]
299mod tests {
300    use super::*;
301    use bytes::Bytes;
302    use gemini_genai_rs::prelude::SessionEvent;
303
304    fn signals() -> SessionSignals {
305        SessionSignals::new(State::new())
306    }
307
308    #[test]
309    fn connected_initializes_state() {
310        let s = signals();
311        s.on_event(&SessionEvent::Connected);
312
313        assert_eq!(s.state.session().get::<u64>("connected_at_ms"), Some(0));
314        assert_eq!(s.state.session().get::<u64>("interrupt_count"), Some(0));
315        assert_eq!(s.state.session().get::<u64>("error_count"), Some(0));
316        assert_eq!(
317            s.state.session().get::<bool>("is_user_speaking"),
318            Some(false)
319        );
320        assert_eq!(
321            s.state.session().get::<bool>("is_model_speaking"),
322            Some(false)
323        );
324        assert_eq!(
325            s.state.session().get::<bool>("go_away_received"),
326            Some(false)
327        );
328        assert_eq!(s.state.session().get::<bool>("resumable"), Some(false));
329        assert_eq!(
330            s.state.session().get::<String>("session_type"),
331            Some("audio_only".to_string())
332        );
333        assert!(s.is_connected.load(Ordering::Relaxed));
334    }
335
336    #[test]
337    fn voice_activity_toggles_user_speaking() {
338        let s = signals();
339        s.on_event(&SessionEvent::Connected);
340        s.on_event(&SessionEvent::VoiceActivityStart);
341        assert_eq!(
342            s.state.session().get::<bool>("is_user_speaking"),
343            Some(true)
344        );
345        s.on_event(&SessionEvent::VoiceActivityEnd);
346        assert_eq!(
347            s.state.session().get::<bool>("is_user_speaking"),
348            Some(false)
349        );
350    }
351
352    #[test]
353    fn interrupted_increments_count() {
354        let s = signals();
355        s.on_event(&SessionEvent::Connected);
356        s.on_event(&SessionEvent::Interrupted);
357        assert_eq!(s.state.session().get::<u64>("interrupt_count"), Some(1));
358        s.on_event(&SessionEvent::Interrupted);
359        assert_eq!(s.state.session().get::<u64>("interrupt_count"), Some(2));
360        s.on_event(&SessionEvent::Interrupted);
361        assert_eq!(s.state.session().get::<u64>("interrupt_count"), Some(3));
362    }
363
364    #[test]
365    fn error_increments_count() {
366        let s = signals();
367        s.on_event(&SessionEvent::Connected);
368        let err = |m: &str| {
369            SessionEvent::Error(gemini_genai_rs::session::SessionError::WebSocket(
370                gemini_genai_rs::session::WebSocketError::ProtocolError(m.into()),
371            ))
372        };
373        s.on_event(&err("oops"));
374        assert_eq!(s.state.session().get::<u64>("error_count"), Some(1));
375        let last: String = s.state.session().get("last_error").unwrap();
376        assert!(last.contains("oops"), "{last}");
377        s.on_event(&err("oops2"));
378        assert_eq!(s.state.session().get::<u64>("error_count"), Some(2));
379        let last: String = s.state.session().get("last_error").unwrap();
380        assert!(last.contains("oops2"), "{last}");
381    }
382
383    #[test]
384    fn phase_changed_sets_model_speaking() {
385        let s = signals();
386        s.on_event(&SessionEvent::Connected);
387        s.on_event(&SessionEvent::PhaseChanged(SessionPhase::ModelSpeaking));
388        assert_eq!(
389            s.state.session().get::<bool>("is_model_speaking"),
390            Some(true)
391        );
392        assert_eq!(
393            s.state.session().get::<String>("phase"),
394            Some("ModelSpeaking".into())
395        );
396        s.on_event(&SessionEvent::PhaseChanged(SessionPhase::Active));
397        assert_eq!(
398            s.state.session().get::<bool>("is_model_speaking"),
399            Some(false)
400        );
401        assert_eq!(
402            s.state.session().get::<String>("phase"),
403            Some("Active".into())
404        );
405    }
406
407    #[test]
408    fn go_away_sets_state() {
409        let s = signals();
410        s.on_event(&SessionEvent::Connected);
411        s.on_event(&SessionEvent::GoAway(Some(std::time::Duration::from_secs(
412            60,
413        ))));
414        assert_eq!(
415            s.state.session().get::<bool>("go_away_received"),
416            Some(true)
417        );
418        assert_eq!(
419            s.state.session().get::<u64>("go_away_time_left_ms"),
420            Some(60_000)
421        );
422        assert!(s.go_away_at.lock().is_some());
423    }
424
425    #[test]
426    fn go_away_without_time_left() {
427        let s = signals();
428        s.on_event(&SessionEvent::Connected);
429        s.on_event(&SessionEvent::GoAway(None));
430        assert_eq!(
431            s.state.session().get::<bool>("go_away_received"),
432            Some(true)
433        );
434        assert_eq!(s.state.session().get::<u64>("go_away_time_left_ms"), None);
435        assert!(s.go_away_at.lock().is_none());
436    }
437
438    #[test]
439    fn session_resume_handle_stored() {
440        let s = signals();
441        s.on_event(&SessionEvent::Connected);
442        s.on_event(&SessionEvent::SessionResumeUpdate(
443            gemini_genai_rs::session::ResumeInfo {
444                handle: "handle-abc".into(),
445                resumable: true,
446                last_consumed_index: None,
447            },
448        ));
449        assert_eq!(s.state.session().get::<bool>("resumable"), Some(true));
450        assert_eq!(s.latest_resume_handle(), Some("handle-abc".to_string()));
451    }
452
453    #[test]
454    fn transcription_stores_last() {
455        let s = signals();
456        s.on_event(&SessionEvent::Connected);
457        s.on_event(&SessionEvent::InputTranscription("hello".into()));
458        assert_eq!(
459            s.state.session().get::<String>("last_input_transcription"),
460            Some("hello".into())
461        );
462        s.on_event(&SessionEvent::OutputTranscription("hi there".into()));
463        assert_eq!(
464            s.state.session().get::<String>("last_output_transcription"),
465            Some("hi there".into())
466        );
467        s.on_event(&SessionEvent::InputTranscription("bye".into()));
468        assert_eq!(
469            s.state.session().get::<String>("last_input_transcription"),
470            Some("bye".into())
471        );
472    }
473
474    #[test]
475    fn session_type_defaults_to_audio_only() {
476        let s = signals();
477        assert_eq!(s.session_type(), SessionType::AudioOnly);
478    }
479
480    #[test]
481    fn mark_video_sent_changes_session_type() {
482        let s = signals();
483        s.on_event(&SessionEvent::Connected);
484        assert_eq!(s.session_type(), SessionType::AudioOnly);
485        s.mark_video_sent();
486        assert_eq!(s.session_type(), SessionType::AudioVideo);
487        assert_eq!(
488            s.state.session().get::<String>("session_type"),
489            Some("audio_video".into())
490        );
491    }
492
493    #[test]
494    fn mark_video_sent_idempotent() {
495        let s = signals();
496        s.on_event(&SessionEvent::Connected);
497        s.mark_video_sent();
498        s.mark_video_sent();
499        assert_eq!(s.session_type(), SessionType::AudioVideo);
500    }
501
502    #[test]
503    fn flush_timing_after_connected() {
504        let s = signals();
505        s.on_event(&SessionEvent::Connected);
506        s.flush_timing();
507        let elapsed: u64 = s.state.session().get("elapsed_ms").unwrap_or(0);
508        assert!(elapsed < 100, "elapsed should be near zero, got {elapsed}");
509        let remaining: u64 = s.state.session().get("remaining_budget_ms").unwrap();
510        let limit = 15 * 60 * 1000u64;
511        assert!(
512            remaining > limit - 1000,
513            "remaining should be near limit, got {remaining}"
514        );
515    }
516
517    #[test]
518    fn flush_timing_respects_video_budget() {
519        let s = signals();
520        s.on_event(&SessionEvent::Connected);
521        s.flush_timing();
522        let remaining_audio: u64 = s.state.session().get("remaining_budget_ms").unwrap();
523        assert!(remaining_audio > 14 * 60 * 1000);
524        s.mark_video_sent();
525        s.flush_timing();
526        let remaining_video: u64 = s.state.session().get("remaining_budget_ms").unwrap();
527        assert!(
528            remaining_video <= 2 * 60 * 1000,
529            "video remaining should be <= 120_000, got {remaining_video}"
530        );
531    }
532
533    #[test]
534    fn latest_resume_handle_initially_none() {
535        let s = signals();
536        assert_eq!(s.latest_resume_handle(), None);
537    }
538
539    #[test]
540    fn latest_resume_handle_updates() {
541        let s = signals();
542        s.on_event(&SessionEvent::SessionResumeUpdate(
543            gemini_genai_rs::session::ResumeInfo {
544                handle: "h1".into(),
545                resumable: true,
546                last_consumed_index: None,
547            },
548        ));
549        assert_eq!(s.latest_resume_handle(), Some("h1".to_string()));
550        s.on_event(&SessionEvent::SessionResumeUpdate(
551            gemini_genai_rs::session::ResumeInfo {
552                handle: "h2".into(),
553                resumable: true,
554                last_consumed_index: Some("5".into()),
555            },
556        ));
557        assert_eq!(s.latest_resume_handle(), Some("h2".to_string()));
558    }
559
560    #[test]
561    fn silence_ms_tracked() {
562        let s = signals();
563        s.on_event(&SessionEvent::Connected);
564        s.flush_timing();
565        let silence: u64 = s.state.session().get("silence_ms").unwrap_or(u64::MAX);
566        assert!(silence < 100, "silence should be near zero, got {silence}");
567    }
568
569    #[test]
570    fn audio_data_updates_activity() {
571        let s = signals();
572        s.on_event(&SessionEvent::Connected);
573        s.on_event(&SessionEvent::AudioData(Bytes::from_static(b"pcm")));
574        s.flush_timing();
575        let silence: u64 = s.state.session().get("silence_ms").unwrap_or(u64::MAX);
576        assert!(silence < 100);
577    }
578
579    #[test]
580    fn turn_complete_updates_activity() {
581        let s = signals();
582        s.on_event(&SessionEvent::Connected);
583        s.on_event(&SessionEvent::TurnComplete);
584        s.flush_timing();
585        let silence: u64 = s.state.session().get("silence_ms").unwrap_or(u64::MAX);
586        assert!(silence < 100);
587    }
588
589    #[test]
590    fn text_complete_updates_activity() {
591        let s = signals();
592        s.on_event(&SessionEvent::Connected);
593        s.on_event(&SessionEvent::TextComplete("done".into()));
594        s.flush_timing();
595        let silence: u64 = s.state.session().get("silence_ms").unwrap_or(u64::MAX);
596        assert!(silence < 100);
597    }
598
599    #[test]
600    fn disconnected_clears_connected_and_sets_flag() {
601        let s = signals();
602        s.on_event(&SessionEvent::Connected);
603        assert!(s.is_connected.load(Ordering::Relaxed));
604        s.on_event(&SessionEvent::Disconnected(Some("server closed".into())));
605        assert!(!s.is_connected.load(Ordering::Relaxed));
606        assert_eq!(s.state.session().get::<bool>("disconnected"), Some(true));
607    }
608
609    #[test]
610    fn disconnected_without_reason() {
611        let s = signals();
612        s.on_event(&SessionEvent::Connected);
613        s.on_event(&SessionEvent::Disconnected(None));
614        assert!(!s.is_connected.load(Ordering::Relaxed));
615        assert_eq!(s.state.session().get::<bool>("disconnected"), Some(true));
616    }
617}