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