1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28pub enum SessionType {
29 AudioOnly,
31 AudioVideo,
33}
34
35pub struct SessionSignals {
49 state: State,
50 start: Instant,
52 clock: SharedClock,
54 connected_at_ns: AtomicU64,
56 is_connected: AtomicBool,
58 last_activity_ns: AtomicU64,
60 has_video: AtomicBool,
62 go_away_at: Mutex<Option<Instant>>,
64 latest_resume_handle: Mutex<Option<String>>,
66}
67
68impl SessionSignals {
69 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 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 }
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 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 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 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 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 pub fn latest_resume_handle(&self) -> Option<String> {
293 self.latest_resume_handle.lock().clone()
294 }
295
296 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}