gemini_genai_rs/vad/
mod.rs

1//! Client-side Voice Activity Detection (VAD).
2//!
3//! WaveKat-backed VAD with the previous dual-threshold energy detector retained
4//! as a fallback for unsupported sample rates or frame sizes. Complements
5//! Gemini's server-side VAD:
6//!
7//! - **Bandwidth savings**: Don't send silence over the network
8//! - **Latency reduction**: Signal `activityStart` before server detects it
9//! - **Barge-in pre-emption**: Flush jitter buffer locally before server confirms
10
11/// VAD configuration parameters.
12#[derive(Debug, Clone)]
13pub struct VadConfig {
14    /// Sample rate in Hz.
15    pub sample_rate: u32,
16    /// Frame duration in milliseconds (typically 10–30ms).
17    pub frame_duration_ms: u32,
18    /// Energy threshold (dBFS) above noise floor to trigger speech start.
19    pub start_threshold_db: f64,
20    /// Energy threshold (dBFS) above noise floor to end speech.
21    pub stop_threshold_db: f64,
22    /// Minimum speech duration in frames before confirming speech.
23    pub min_speech_frames: u32,
24    /// Hangover duration in frames — keeps "speaking" state after energy drops.
25    pub hangover_frames: u32,
26    /// ZCR range for speech confirmation (low, high).
27    pub speech_zcr_range: (f64, f64),
28    /// Initial noise floor estimate (dBFS).
29    pub initial_noise_floor_db: f64,
30    /// Number of pre-speech frames to buffer.
31    pub pre_speech_frames: usize,
32}
33
34impl Default for VadConfig {
35    fn default() -> Self {
36        Self {
37            sample_rate: 16000,
38            frame_duration_ms: 30,
39            start_threshold_db: 15.0,
40            stop_threshold_db: 10.0,
41            min_speech_frames: 3,
42            hangover_frames: 10,
43            speech_zcr_range: (0.02, 0.5),
44            initial_noise_floor_db: -60.0,
45            pre_speech_frames: 3,
46        }
47    }
48}
49
50impl VadConfig {
51    /// Number of samples per frame.
52    pub fn frame_size(&self) -> usize {
53        (self.sample_rate * self.frame_duration_ms / 1000) as usize
54    }
55
56    /// Tuned for noisy environments **behind a denoiser** (the L2
57    /// `voice::Denoiser`, feature `denoise`) — do not use on a raw noisy
58    /// stream, where the energy floor latches regardless of threshold.
59    ///
60    /// Found by a latency-aware parameter sweep over labeled street-traffic,
61    /// pink, and white scenes (clean → 0 dB SNR): with the signal denoised,
62    /// the threshold can be *raised* to 21 dB (rejecting horn and engine
63    /// residue) while the onset confirmation drops to a single frame
64    /// (~150–310 ms measured onset). On the same benchmark the shipped
65    /// defaults behind the denoiser score 174–312 ms onset with more
66    /// residual open time; on street traffic at 0 dB this preset reads
67    /// 0 false activations / 3 of 3 utterances / 0 % stuck-open, and it
68    /// held 0 false client activations in a live Gemini session streaming
69    /// 0 dB traffic. See the hardening chapter for the full study.
70    pub fn noisy_street() -> Self {
71        Self {
72            start_threshold_db: 21.0,
73            stop_threshold_db: 16.0,
74            min_speech_frames: 1,
75            hangover_frames: 10,
76            ..Self::default()
77        }
78    }
79}
80
81/// VAD state machine states.
82#[derive(Debug, Clone, Copy, PartialEq, Eq)]
83pub enum VadState {
84    /// No speech detected.
85    Silence,
86    /// Energy exceeded threshold but min duration not yet met.
87    PendingSpeech,
88    /// Speech confirmed.
89    Speech,
90    /// Energy dropped but still in hangover period.
91    Hangover,
92}
93
94/// Events emitted by the VAD.
95#[derive(Debug, Clone, Copy, PartialEq, Eq)]
96pub enum VadEvent {
97    /// Speech onset detected.
98    SpeechStart,
99    /// Speech ended (after hangover).
100    SpeechEnd,
101}
102
103/// Voice Activity Detector with adaptive noise floor.
104pub struct VoiceActivityDetector {
105    config: VadConfig,
106    #[cfg(feature = "vad-wavekat")]
107    wavekat: Option<WaveKatWebRtcBackend>,
108    use_wavekat: bool,
109    state: VadState,
110    /// Adaptive noise floor estimate (dBFS).
111    noise_floor_db: f64,
112    /// Frames spent in current state.
113    state_frames: u32,
114    /// Number of frames used for noise adaptation.
115    noise_adapt_frames: u64,
116    /// Last backend speech probability/decision, normalized to 0.0..1.0.
117    last_probability: Option<f32>,
118    /// Circular buffer of pre-speech frames.
119    pre_speech_buf: Vec<Vec<i16>>,
120    pre_speech_idx: usize,
121}
122
123impl VoiceActivityDetector {
124    /// Create a new VAD with the given configuration.
125    pub fn new(config: VadConfig) -> Self {
126        Self::new_with_backend(config, true)
127    }
128
129    fn new_with_backend(config: VadConfig, prefer_wavekat: bool) -> Self {
130        let frame_size = config.frame_size();
131        let pre_speech_buf: Vec<Vec<i16>> = (0..config.pre_speech_frames)
132            .map(|_| vec![0i16; frame_size])
133            .collect();
134        #[cfg(feature = "vad-wavekat")]
135        let wavekat = prefer_wavekat
136            .then(|| WaveKatWebRtcBackend::new(&config))
137            .flatten();
138        #[cfg(feature = "vad-wavekat")]
139        let use_wavekat = wavekat.is_some();
140        #[cfg(not(feature = "vad-wavekat"))]
141        let use_wavekat = {
142            let _ = prefer_wavekat;
143            false
144        };
145        Self {
146            noise_floor_db: config.initial_noise_floor_db,
147            state: VadState::Silence,
148            state_frames: 0,
149            noise_adapt_frames: 0,
150            last_probability: None,
151            pre_speech_buf,
152            pre_speech_idx: 0,
153            #[cfg(feature = "vad-wavekat")]
154            wavekat,
155            use_wavekat,
156            config,
157        }
158    }
159
160    #[cfg(test)]
161    fn new_energy(config: VadConfig) -> Self {
162        Self::new_with_backend(config, false)
163    }
164
165    /// Current VAD state.
166    pub fn state(&self) -> VadState {
167        self.state
168    }
169
170    /// Whether speech is currently detected (Speech or Hangover state).
171    pub fn is_speaking(&self) -> bool {
172        matches!(self.state, VadState::Speech | VadState::Hangover)
173    }
174
175    /// Current noise floor estimate (dBFS).
176    pub fn noise_floor_db(&self) -> f64 {
177        self.noise_floor_db
178    }
179
180    /// Whether this detector is currently using the WaveKat backend.
181    pub fn is_wavekat_backed(&self) -> bool {
182        self.use_wavekat
183    }
184
185    /// Name of the active backend.
186    pub fn backend_name(&self) -> &'static str {
187        if self.use_wavekat {
188            "wavekat-webrtc"
189        } else {
190            "energy-zcr"
191        }
192    }
193
194    /// Last normalized speech probability or binary backend decision.
195    pub fn last_probability(&self) -> Option<f32> {
196        self.last_probability
197    }
198
199    /// Get pre-speech frames (the frames captured just before speech onset).
200    pub fn drain_pre_speech(&mut self) -> Vec<Vec<i16>> {
201        let frame_size = self.config.frame_size();
202        let mut fresh: Vec<Vec<i16>> = (0..self.config.pre_speech_frames)
203            .map(|_| vec![0i16; frame_size])
204            .collect();
205        std::mem::swap(&mut self.pre_speech_buf, &mut fresh);
206        self.pre_speech_idx = 0;
207        fresh
208    }
209
210    /// Process a single audio frame and return any state-change event.
211    pub fn process_frame(&mut self, samples: &[i16]) -> Option<VadEvent> {
212        let energy_db = compute_energy_db(samples);
213        let zcr = compute_zcr(samples);
214        let energy_above_noise = energy_db - self.noise_floor_db;
215
216        let wavekat_decision = self.wavekat_decision(samples);
217        let energy_speech_like = energy_above_noise > self.config.start_threshold_db
218            && zcr >= self.config.speech_zcr_range.0
219            && zcr <= self.config.speech_zcr_range.1;
220        let energy_above_stop = energy_above_noise > self.config.stop_threshold_db;
221        let is_speech_like = wavekat_decision.unwrap_or(energy_speech_like);
222        let is_above_stop = wavekat_decision.unwrap_or(energy_above_stop);
223        if wavekat_decision.is_none() {
224            self.last_probability = Some(if energy_speech_like { 1.0 } else { 0.0 });
225        }
226
227        match self.state {
228            VadState::Silence => {
229                // Update noise floor during confirmed silence
230                self.update_noise_floor(energy_db);
231
232                // Store pre-speech frame (copy into pre-allocated slot, zero-alloc)
233                if self.config.pre_speech_frames > 0 && !self.pre_speech_buf.is_empty() {
234                    let idx = self.pre_speech_idx % self.config.pre_speech_frames;
235                    let buf = &mut self.pre_speech_buf[idx];
236                    buf.resize(samples.len(), 0);
237                    buf.copy_from_slice(samples);
238                    self.pre_speech_idx += 1;
239                }
240
241                if is_speech_like {
242                    self.state_frames = 1;
243                    // Check if this single frame meets the minimum requirement
244                    if self.state_frames >= self.config.min_speech_frames {
245                        self.state = VadState::Speech;
246                        self.state_frames = 0;
247                        return Some(VadEvent::SpeechStart);
248                    }
249                    self.state = VadState::PendingSpeech;
250                }
251                None
252            }
253
254            VadState::PendingSpeech => {
255                if is_speech_like {
256                    self.state_frames += 1;
257                    if self.state_frames >= self.config.min_speech_frames {
258                        self.state = VadState::Speech;
259                        self.state_frames = 0;
260                        Some(VadEvent::SpeechStart)
261                    } else {
262                        None
263                    }
264                } else {
265                    // False alarm — go back to silence
266                    self.state = VadState::Silence;
267                    self.state_frames = 0;
268                    None
269                }
270            }
271
272            VadState::Speech => {
273                if !is_above_stop {
274                    self.state = VadState::Hangover;
275                    self.state_frames = 1;
276                }
277                None
278            }
279
280            VadState::Hangover => {
281                if is_above_stop {
282                    // Speech resumed — back to Speech
283                    self.state = VadState::Speech;
284                    self.state_frames = 0;
285                    None
286                } else {
287                    self.state_frames += 1;
288                    if self.state_frames >= self.config.hangover_frames {
289                        self.state = VadState::Silence;
290                        self.state_frames = 0;
291                        for buf in &mut self.pre_speech_buf {
292                            buf.iter_mut().for_each(|s| *s = 0);
293                        }
294                        self.pre_speech_idx = 0;
295                        Some(VadEvent::SpeechEnd)
296                    } else {
297                        None
298                    }
299                }
300            }
301        }
302    }
303
304    /// Update the adaptive noise floor using EWMA.
305    fn update_noise_floor(&mut self, energy_db: f64) {
306        self.noise_adapt_frames += 1;
307        // Alpha decreases over time: fast initial adaptation, slow drift
308        let alpha = 0.01_f64.min(1.0 / self.noise_adapt_frames as f64);
309        self.noise_floor_db = self.noise_floor_db * (1.0 - alpha) + energy_db * alpha;
310    }
311
312    #[allow(
313        clippy::unused_self,
314        reason = "self carries the backend and last probability once `vad-wavekat` is compiled in; \
315                  the feature-off body is a stub behind the same call site"
316    )]
317    fn wavekat_decision(&mut self, samples: &[i16]) -> Option<bool> {
318        #[cfg(feature = "vad-wavekat")]
319        {
320            let probability = self
321                .wavekat
322                .as_mut()
323                .and_then(|backend| backend.process(samples, self.config.sample_rate));
324            self.last_probability = probability;
325            probability.map(|probability| probability >= 0.5)
326        }
327        #[cfg(not(feature = "vad-wavekat"))]
328        {
329            let _ = samples;
330            None
331        }
332    }
333
334    /// Reset the VAD to its initial state.
335    pub fn reset(&mut self) {
336        self.state = VadState::Silence;
337        self.state_frames = 0;
338        self.noise_adapt_frames = 0;
339        self.last_probability = None;
340        self.noise_floor_db = self.config.initial_noise_floor_db;
341        for buf in &mut self.pre_speech_buf {
342            buf.iter_mut().for_each(|s| *s = 0);
343        }
344        self.pre_speech_idx = 0;
345    }
346}
347
348#[cfg(feature = "vad-wavekat")]
349struct WaveKatWebRtcBackend {
350    detector: wavekat_vad::backends::webrtc::WebRtcVad,
351}
352
353#[cfg(feature = "vad-wavekat")]
354impl WaveKatWebRtcBackend {
355    fn new(config: &VadConfig) -> Option<Self> {
356        if !matches!(config.sample_rate, 8000 | 16000 | 32000 | 48000) {
357            return None;
358        }
359        if !matches!(config.frame_duration_ms, 10 | 20 | 30) {
360            return None;
361        }
362
363        let detector = wavekat_vad::backends::webrtc::WebRtcVad::with_frame_duration(
364            config.sample_rate,
365            wavekat_vad::backends::webrtc::WebRtcVadMode::Aggressive,
366            config.frame_duration_ms,
367        )
368        .ok()?;
369        Some(Self { detector })
370    }
371
372    fn process(&mut self, samples: &[i16], sample_rate: u32) -> Option<f32> {
373        use wavekat_vad::VoiceActivityDetector as _;
374
375        self.detector.process(samples, sample_rate).ok()
376    }
377}
378
379/// Compute RMS energy in dBFS for a frame of PCM16 samples.
380fn compute_energy_db(samples: &[i16]) -> f64 {
381    if samples.is_empty() {
382        return -96.0;
383    }
384
385    let sum_sq: f64 = samples.iter().map(|&s| (s as f64) * (s as f64)).sum();
386    let rms = (sum_sq / samples.len() as f64).sqrt();
387    let db = 20.0 * (rms / 32767.0).log10();
388    db.max(-96.0) // Floor at -96 dBFS
389}
390
391/// Compute zero-crossing rate for a frame of PCM16 samples.
392fn compute_zcr(samples: &[i16]) -> f64 {
393    if samples.len() < 2 {
394        return 0.0;
395    }
396
397    let crossings = samples
398        .windows(2)
399        .filter(|w| (w[0] >= 0) != (w[1] >= 0))
400        .count();
401
402    crossings as f64 / (samples.len() - 1) as f64
403}
404
405#[cfg(test)]
406mod tests {
407    use super::*;
408
409    fn make_vad() -> VoiceActivityDetector {
410        VoiceActivityDetector::new_energy(VadConfig {
411            sample_rate: 16000,
412            frame_duration_ms: 20,
413            start_threshold_db: 15.0,
414            stop_threshold_db: 10.0,
415            min_speech_frames: 2,
416            hangover_frames: 3,
417            speech_zcr_range: (0.01, 0.9),
418            initial_noise_floor_db: -60.0,
419            pre_speech_frames: 2,
420        })
421    }
422
423    fn silence_frame(len: usize) -> Vec<i16> {
424        vec![0i16; len]
425    }
426
427    fn speech_frame(len: usize, amplitude: i16) -> Vec<i16> {
428        // Generate a simple alternating signal that has both energy and ZCR
429        (0..len)
430            .map(|i| if i % 4 < 2 { amplitude } else { -amplitude })
431            .collect()
432    }
433
434    #[test]
435    fn starts_silent() {
436        let vad = make_vad();
437        assert_eq!(vad.state(), VadState::Silence);
438        assert!(!vad.is_speaking());
439    }
440
441    #[cfg(feature = "vad-wavekat")]
442    #[test]
443    fn default_detector_uses_wavekat_for_supported_frames() {
444        let vad = VoiceActivityDetector::new(VadConfig {
445            sample_rate: 16000,
446            frame_duration_ms: 20,
447            ..VadConfig::default()
448        });
449        assert!(vad.is_wavekat_backed());
450    }
451
452    #[cfg(feature = "vad-wavekat")]
453    #[test]
454    fn unsupported_frames_fall_back_to_energy_detector() {
455        let vad = VoiceActivityDetector::new(VadConfig {
456            sample_rate: 16000,
457            frame_duration_ms: 32,
458            ..VadConfig::default()
459        });
460        assert!(!vad.is_wavekat_backed());
461    }
462
463    #[test]
464    fn silence_stays_silent() {
465        let mut vad = make_vad();
466        let frame = silence_frame(320);
467        for _ in 0..10 {
468            let event = vad.process_frame(&frame);
469            assert!(event.is_none());
470        }
471        assert_eq!(vad.state(), VadState::Silence);
472    }
473
474    #[test]
475    fn speech_detected_after_min_frames() {
476        let mut vad = make_vad();
477        let frame = speech_frame(320, 10000);
478
479        // Frame 1: PendingSpeech
480        let e1 = vad.process_frame(&frame);
481        assert!(e1.is_none());
482        assert_eq!(vad.state(), VadState::PendingSpeech);
483
484        // Frame 2: min_speech_frames = 2 → SpeechStart
485        let e2 = vad.process_frame(&frame);
486        assert_eq!(e2, Some(VadEvent::SpeechStart));
487        assert_eq!(vad.state(), VadState::Speech);
488        assert!(vad.is_speaking());
489    }
490
491    #[test]
492    fn speech_end_after_hangover() {
493        let mut vad = make_vad();
494        let speech = speech_frame(320, 10000);
495        let silence = silence_frame(320);
496
497        // Trigger speech
498        vad.process_frame(&speech);
499        vad.process_frame(&speech);
500        assert_eq!(vad.state(), VadState::Speech);
501
502        // Drop energy → hangover
503        vad.process_frame(&silence);
504        assert_eq!(vad.state(), VadState::Hangover);
505
506        // Hangover frames 2 and 3
507        vad.process_frame(&silence);
508        let e = vad.process_frame(&silence);
509        assert_eq!(e, Some(VadEvent::SpeechEnd));
510        assert_eq!(vad.state(), VadState::Silence);
511    }
512
513    #[test]
514    fn speech_resumes_during_hangover() {
515        let mut vad = make_vad();
516        let speech = speech_frame(320, 10000);
517        let silence = silence_frame(320);
518
519        // Trigger speech
520        vad.process_frame(&speech);
521        vad.process_frame(&speech);
522        assert_eq!(vad.state(), VadState::Speech);
523
524        // Brief silence → hangover
525        vad.process_frame(&silence);
526        assert_eq!(vad.state(), VadState::Hangover);
527
528        // Speech resumes
529        let e = vad.process_frame(&speech);
530        assert!(e.is_none()); // No event, just resumes
531        assert_eq!(vad.state(), VadState::Speech);
532    }
533
534    #[test]
535    fn false_alarm_returns_to_silence() {
536        let mut vad = make_vad();
537        let speech = speech_frame(320, 10000);
538        let silence = silence_frame(320);
539
540        // 1 speech frame → PendingSpeech
541        vad.process_frame(&speech);
542        assert_eq!(vad.state(), VadState::PendingSpeech);
543
544        // Then silence → back to Silence (false alarm)
545        vad.process_frame(&silence);
546        assert_eq!(vad.state(), VadState::Silence);
547    }
548
549    #[test]
550    fn energy_db_calculation() {
551        // Full-scale sine approximation
552        let full_scale: Vec<i16> = (0..320).map(|_| i16::MAX).collect();
553        let db = compute_energy_db(&full_scale);
554        assert!(db > -1.0); // Should be near 0 dBFS
555
556        let silence = vec![0i16; 320];
557        let db_silence = compute_energy_db(&silence);
558        assert_eq!(db_silence, -96.0);
559    }
560
561    #[test]
562    fn zcr_calculation() {
563        // Alternating signal → high ZCR
564        let alternating: Vec<i16> = (0..100)
565            .map(|i| if i % 2 == 0 { 1000 } else { -1000 })
566            .collect();
567        let zcr = compute_zcr(&alternating);
568        assert!(zcr > 0.9);
569
570        // Constant signal → zero ZCR
571        let constant = vec![1000i16; 100];
572        let zcr_const = compute_zcr(&constant);
573        assert_eq!(zcr_const, 0.0);
574    }
575
576    #[test]
577    fn noise_floor_adapts() {
578        let mut vad = make_vad();
579        // Feed low-energy frames — noise floor should move toward them
580        let low_noise: Vec<i16> = vec![10; 320]; // Very quiet
581        for _ in 0..100 {
582            vad.process_frame(&low_noise);
583        }
584        // Noise floor should have adapted upward from -60 dBFS
585        assert!(vad.noise_floor_db() > -96.0);
586    }
587
588    #[test]
589    fn reset_clears_state() {
590        let mut vad = make_vad();
591        let speech = speech_frame(320, 10000);
592        vad.process_frame(&speech);
593        vad.process_frame(&speech);
594        assert_eq!(vad.state(), VadState::Speech);
595
596        vad.reset();
597        assert_eq!(vad.state(), VadState::Silence);
598        assert_eq!(vad.noise_floor_db(), -60.0);
599    }
600
601    #[test]
602    fn noisy_street_preset_is_stricter_and_faster_than_default() {
603        let preset = VadConfig::noisy_street();
604        let default = VadConfig::default();
605        // Higher bar to open (noise rejection) with hysteresis preserved…
606        assert!(preset.start_threshold_db > default.start_threshold_db);
607        assert!(preset.stop_threshold_db < preset.start_threshold_db);
608        // …and a shorter onset confirmation (lower barge-in latency).
609        assert!(preset.min_speech_frames < default.min_speech_frames);
610        assert!(preset.min_speech_frames >= 1);
611        // Frame geometry unchanged — presets tune sensitivity, not timing base.
612        assert_eq!(preset.frame_size(), default.frame_size());
613    }
614
615    #[test]
616    fn pending_speech_emits_on_first_frame_when_min_speech_frames_one() {
617        let mut vad = VoiceActivityDetector::new_energy(VadConfig {
618            sample_rate: 16000,
619            frame_duration_ms: 20,
620            start_threshold_db: 15.0,
621            stop_threshold_db: 10.0,
622            min_speech_frames: 1,
623            hangover_frames: 3,
624            speech_zcr_range: (0.01, 0.9),
625            initial_noise_floor_db: -60.0,
626            pre_speech_frames: 2,
627        });
628
629        let frame = speech_frame(320, 10000);
630
631        // With min_speech_frames=1, should emit SpeechStart on frame 1
632        let e = vad.process_frame(&frame);
633        assert_eq!(
634            e,
635            Some(VadEvent::SpeechStart),
636            "Expected SpeechStart on first speech frame when min_speech_frames=1"
637        );
638        assert_eq!(vad.state(), VadState::Speech);
639    }
640
641    #[test]
642    fn hangover_exits_after_min_frames_when_hangover_frames_one() {
643        let mut vad = VoiceActivityDetector::new_energy(VadConfig {
644            sample_rate: 16000,
645            frame_duration_ms: 20,
646            start_threshold_db: 15.0,
647            stop_threshold_db: 10.0,
648            min_speech_frames: 2,
649            hangover_frames: 1,
650            speech_zcr_range: (0.01, 0.9),
651            initial_noise_floor_db: -60.0,
652            pre_speech_frames: 2,
653        });
654
655        let speech = speech_frame(320, 10000);
656        let silence = silence_frame(320);
657
658        // Trigger speech (need 2 frames with min_speech_frames=2)
659        vad.process_frame(&speech);
660        assert_eq!(vad.state(), VadState::PendingSpeech);
661        vad.process_frame(&speech);
662        assert_eq!(vad.state(), VadState::Speech);
663
664        // Drop energy -> hangover (frame 1 in hangover)
665        let e1 = vad.process_frame(&silence);
666        assert!(
667            e1.is_none(),
668            "Should not emit SpeechEnd when entering hangover"
669        );
670        assert_eq!(vad.state(), VadState::Hangover);
671
672        // With hangover_frames=1, should exit hangover on frame 1 (the entry)
673        // But current implementation exits on frame 2
674        let e2 = vad.process_frame(&silence);
675        assert_eq!(
676            e2,
677            Some(VadEvent::SpeechEnd),
678            "Expected SpeechEnd after 1 hangover frame when hangover_frames=1"
679        );
680        assert_eq!(vad.state(), VadState::Silence);
681    }
682}