1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum SessionType {
27 AudioOnly,
29 AudioVideo,
31}
32
33pub struct SessionSignals {
47 state: State,
48 start: Instant,
50 connected_at_ns: AtomicU64,
52 is_connected: AtomicBool,
54 last_activity_ns: AtomicU64,
56 has_video: AtomicBool,
58 go_away_at: Mutex<Option<Instant>>,
60 latest_resume_handle: Mutex<Option<String>>,
62}
63
64impl SessionSignals {
65 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 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 }
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 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 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 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 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 pub fn latest_resume_handle(&self) -> Option<String> {
287 self.latest_resume_handle.lock().clone()
288 }
289
290 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}