diff --git a/CHANGELOG.md b/CHANGELOG.md index b4d742b..0bd5be9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,6 +16,17 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - NVIDIA Parakeet support is available in the bundled C library but is not yet exposed through the Rust API. - The new upstream VAD segment and VAD-mapped token timestamp accessors are not exposed: they are only populated by upstream's built-in `whisper_full` VAD, which does not run for per-state transcription (`whisper_full_with_state`, used by this crate; see ggml-org/whisper.cpp#3423). Use the crate's own VAD pipeline instead. +### Added + +- `WhisperState::full_get_segment_no_speech_prob()`, previously only used internally by the temperature-fallback transcriber. + +### Fixed + +- **Breaking:** segment timestamps are now real milliseconds. whisper.cpp reports segment times in centiseconds, and the crate previously passed them through unconverted, so `Segment::start_ms`/`end_ms`, `start_seconds()`/`end_seconds()`, `WhisperState::full_get_segment_timestamps()`, and the `start`/`end` values passed to `WhisperStreamPcm::run` callbacks were 10x too small. This affects `transcribe*`, `WhisperStream`, `WhisperStreamPcm`, and the temperature-fallback transcriber. The raw `whisper_token_data` returned by `full_get_token_data()` is unchanged and documented as centiseconds. +- Fixed a use-after-free in `FullParams::suppress_regex()`: the regex string was freed immediately after being set, so whisper.cpp read freed memory during transcription. +- Fixed `FullParams::prompt_tokens()` storing a borrowed pointer that dangled once the caller's slice was dropped or the params were moved or cloned. The tokens are now copied into the params. +- **Breaking:** `WhisperState` result getters now validate segment and token indices instead of passing them to whisper.cpp, which does not bounds-check (out-of-range indices were undefined behaviour). `full_get_segment_text()` and `full_get_token_text()` return `WhisperError::InvalidParameter`, `full_get_token_data()` returns `None`, and the plain-value getters (`full_get_segment_timestamps()`, `full_get_segment_speaker_turn_next()`, `full_n_tokens()`, `full_get_token_id()`, `full_get_token_prob()`) panic, like slice indexing. + ## [0.1.5] - 2026-06-12 ### Added diff --git a/whisper-cpp-plus/src/enhanced/fallback.rs b/whisper-cpp-plus/src/enhanced/fallback.rs index 8ba3efb..7b41539 100644 --- a/whisper-cpp-plus/src/enhanced/fallback.rs +++ b/whisper-cpp-plus/src/enhanced/fallback.rs @@ -6,7 +6,6 @@ use crate::{FullParams, Result, Segment, TranscriptionResult, WhisperError, Whis use flate2::write::ZlibEncoder; use flate2::Compression; use std::io::Write; -use whisper_cpp_plus_sys as ffi; /// Quality thresholds for transcription validation #[derive(Debug, Clone)] @@ -166,10 +165,7 @@ impl<'a> EnhancedWhisperState<'a> { /// Get no-speech probability for a segment (enhanced feature) fn get_no_speech_prob(&self, segment_idx: i32) -> f32 { - unsafe { - // Direct FFI call using the exposed ptr - ffi::whisper_full_get_segment_no_speech_prob_from_state(self.state.ptr, segment_idx) - } + self.state.full_get_segment_no_speech_prob(segment_idx) } /// Calculate average log probability from token probabilities diff --git a/whisper-cpp-plus/src/params.rs b/whisper-cpp-plus/src/params.rs index 76eabab..62add32 100644 --- a/whisper-cpp-plus/src/params.rs +++ b/whisper-cpp-plus/src/params.rs @@ -10,8 +10,12 @@ pub enum SamplingStrategy { #[derive(Clone)] pub struct FullParams { pub(crate) inner: ffi::whisper_full_params, + // Data referenced by pointer fields in `inner` is owned here and wired up in `as_raw`, so + // the pointers stay valid across moves and clones. language: Option, initial_prompt: Option, + suppress_regex: Option, + prompt_tokens: Vec, } // FullParams is Send and Sync because we only use it in controlled contexts @@ -43,6 +47,8 @@ impl FullParams { inner, language: None, initial_prompt: None, + suppress_regex: None, + prompt_tokens: Vec::new(), }; params.inner.n_threads = (num_cpus::get() / 2).max(1) as i32; @@ -66,14 +72,24 @@ impl FullParams { params.initial_prompt = prompt.as_ptr(); } + params.suppress_regex = self + .suppress_regex + .as_ref() + .map_or(std::ptr::null(), |regex| regex.as_ptr()); + + if self.prompt_tokens.is_empty() { + params.prompt_tokens = std::ptr::null(); + params.prompt_n_tokens = 0; + } else { + params.prompt_tokens = self.prompt_tokens.as_ptr(); + params.prompt_n_tokens = self.prompt_tokens.len() as i32; + } + params } pub fn language(mut self, lang: &str) -> Self { self.language = CString::new(lang).ok(); - if let Some(ref lang_cstr) = self.language { - self.inner.language = lang_cstr.as_ptr(); - } self } @@ -162,28 +178,29 @@ impl FullParams { self } + /// Suppress tokens matching `suppress_regex`; `None` clears it. + /// + /// A regex containing an interior NUL byte cannot be passed to C and is ignored. pub fn suppress_regex(mut self, suppress_regex: Option<&str>) -> Self { - if let Some(regex) = suppress_regex { - if let Ok(c_regex) = CString::new(regex) { - self.inner.suppress_regex = c_regex.as_ptr(); + match suppress_regex { + Some(regex) => { + if let Ok(c_regex) = CString::new(regex) { + self.suppress_regex = Some(c_regex); + } } - } else { - self.inner.suppress_regex = std::ptr::null(); + None => self.suppress_regex = None, } self } pub fn initial_prompt(mut self, prompt: &str) -> Self { self.initial_prompt = CString::new(prompt).ok(); - if let Some(ref prompt_cstr) = self.initial_prompt { - self.inner.initial_prompt = prompt_cstr.as_ptr(); - } self } + /// Prompt the decoder with these tokens. The tokens are copied. pub fn prompt_tokens(mut self, tokens: &[i32]) -> Self { - self.inner.prompt_tokens = tokens.as_ptr(); - self.inner.prompt_n_tokens = tokens.len() as i32; + self.prompt_tokens = tokens.to_vec(); self } @@ -313,3 +330,60 @@ impl Default for TranscriptionParamsBuilder { Self::new() } } + +#[cfg(test)] +mod tests { + use super::*; + use std::ffi::CStr; + + fn raw_str(ptr: *const std::os::raw::c_char) -> String { + assert!(!ptr.is_null()); + unsafe { CStr::from_ptr(ptr) }.to_str().unwrap().to_owned() + } + + #[test] + fn string_params_survive_clone_and_drop() { + let original = FullParams::default() + .language("de") + .initial_prompt("hello") + .suppress_regex(Some("[0-9]+")); + let cloned = original.clone(); + drop(original); + + let raw = cloned.as_raw(); + assert_eq!(raw_str(raw.language), "de"); + assert_eq!(raw_str(raw.initial_prompt), "hello"); + assert_eq!(raw_str(raw.suppress_regex), "[0-9]+"); + } + + #[test] + fn suppress_regex_none_clears() { + let params = FullParams::default() + .suppress_regex(Some("abc")) + .suppress_regex(None); + assert!(params.as_raw().suppress_regex.is_null()); + assert!(FullParams::default().as_raw().suppress_regex.is_null()); + } + + #[test] + fn prompt_tokens_are_owned() { + let params = { + let tokens = vec![50257, 1, 2, 3]; + FullParams::default().prompt_tokens(&tokens) + }; + let cloned = params.clone(); + drop(params); + + let raw = cloned.as_raw(); + assert_eq!(raw.prompt_n_tokens, 4); + let tokens = unsafe { std::slice::from_raw_parts(raw.prompt_tokens, 4) }; + assert_eq!(tokens, &[50257, 1, 2, 3]); + } + + #[test] + fn empty_prompt_tokens_are_null() { + let raw = FullParams::default().prompt_tokens(&[]).as_raw(); + assert!(raw.prompt_tokens.is_null()); + assert_eq!(raw.prompt_n_tokens, 0); + } +} diff --git a/whisper-cpp-plus/src/state.rs b/whisper-cpp-plus/src/state.rs index e0f93f4..a95bf11 100644 --- a/whisper-cpp-plus/src/state.rs +++ b/whisper-cpp-plus/src/state.rs @@ -102,7 +102,48 @@ impl WhisperState { unsafe { ffi::whisper_full_lang_id_from_state(self.ptr) } } + // The whisper.cpp result getters index their vectors without bounds checks, so every + // wrapper validates indices before calling into C. + + fn segment_in_range(&self, i_segment: i32) -> bool { + i_segment >= 0 && i_segment < self.full_n_segments() + } + + fn token_in_range(&self, i_segment: i32, i_token: i32) -> bool { + self.segment_in_range(i_segment) + && i_token >= 0 + && i_token < unsafe { ffi::whisper_full_n_tokens_from_state(self.ptr, i_segment) } + } + + fn assert_segment_in_range(&self, i_segment: i32) { + assert!( + self.segment_in_range(i_segment), + "segment index {} out of range (n_segments = {})", + i_segment, + self.full_n_segments() + ); + } + + fn assert_token_in_range(&self, i_segment: i32, i_token: i32) { + assert!( + self.token_in_range(i_segment, i_token), + "token index ({}, {}) out of range", + i_segment, + i_token + ); + } + + /// Returns the text of segment `i_segment`. + /// + /// Returns [`WhisperError::InvalidParameter`] if the index is out of range. pub fn full_get_segment_text(&self, i_segment: i32) -> Result { + if !self.segment_in_range(i_segment) { + return Err(WhisperError::InvalidParameter(format!( + "segment index {} out of range", + i_segment + ))); + } + let text_ptr = unsafe { ffi::whisper_full_get_segment_text_from_state(self.ptr, i_segment) }; @@ -114,23 +155,62 @@ impl WhisperState { Ok(c_str.to_string_lossy().into_owned()) } + /// Returns the `(start, end)` time of segment `i_segment` in milliseconds. + /// + /// # Panics + /// + /// Panics if `i_segment` is out of range. pub fn full_get_segment_timestamps(&self, i_segment: i32) -> (i64, i64) { + self.assert_segment_in_range(i_segment); + // whisper.cpp reports segment times in centiseconds (10 ms units). unsafe { let t0 = ffi::whisper_full_get_segment_t0_from_state(self.ptr, i_segment); let t1 = ffi::whisper_full_get_segment_t1_from_state(self.ptr, i_segment); - (t0, t1) + (t0 * 10, t1 * 10) } } + /// Returns whether the next segment starts with a speaker turn (tinydiarize). + /// + /// # Panics + /// + /// Panics if `i_segment` is out of range. pub fn full_get_segment_speaker_turn_next(&self, i_segment: i32) -> bool { + self.assert_segment_in_range(i_segment); unsafe { ffi::whisper_full_get_segment_speaker_turn_next_from_state(self.ptr, i_segment) } } + /// Returns the no-speech probability of segment `i_segment`. + /// + /// # Panics + /// + /// Panics if `i_segment` is out of range. + pub fn full_get_segment_no_speech_prob(&self, i_segment: i32) -> f32 { + self.assert_segment_in_range(i_segment); + unsafe { ffi::whisper_full_get_segment_no_speech_prob_from_state(self.ptr, i_segment) } + } + + /// Returns the number of tokens in segment `i_segment`. + /// + /// # Panics + /// + /// Panics if `i_segment` is out of range. pub fn full_n_tokens(&self, i_segment: i32) -> i32 { + self.assert_segment_in_range(i_segment); unsafe { ffi::whisper_full_n_tokens_from_state(self.ptr, i_segment) } } + /// Returns the text of token `i_token` in segment `i_segment`. + /// + /// Returns [`WhisperError::InvalidParameter`] if either index is out of range. pub fn full_get_token_text(&self, i_segment: i32, i_token: i32) -> Result { + if !self.token_in_range(i_segment, i_token) { + return Err(WhisperError::InvalidParameter(format!( + "token index ({}, {}) out of range", + i_segment, i_token + ))); + } + let text_ptr = unsafe { ffi::whisper_full_get_token_text_from_state( self._context.0, @@ -148,26 +228,40 @@ impl WhisperState { Ok(c_str.to_string_lossy().into_owned()) } + /// Returns the id of token `i_token` in segment `i_segment`. + /// + /// # Panics + /// + /// Panics if either index is out of range. pub fn full_get_token_id(&self, i_segment: i32, i_token: i32) -> i32 { + self.assert_token_in_range(i_segment, i_token); unsafe { ffi::whisper_full_get_token_id_from_state(self.ptr, i_segment, i_token) } } + /// Returns the raw token data for token `i_token` in segment `i_segment`, or `None` if + /// either index is out of range. + /// + /// This is the unmodified whisper.cpp struct: its `t0`, `t1` and `t_dtw` fields are in + /// centiseconds (10 ms units), unlike [`WhisperState::full_get_segment_timestamps`]. pub fn full_get_token_data( &self, i_segment: i32, i_token: i32, ) -> Option { - let data = - unsafe { ffi::whisper_full_get_token_data_from_state(self.ptr, i_segment, i_token) }; - - if data.id == -1 { - None - } else { - Some(data) + if !self.token_in_range(i_segment, i_token) { + return None; } + + Some(unsafe { ffi::whisper_full_get_token_data_from_state(self.ptr, i_segment, i_token) }) } + /// Returns the probability of token `i_token` in segment `i_segment`. + /// + /// # Panics + /// + /// Panics if either index is out of range. pub fn full_get_token_prob(&self, i_segment: i32, i_token: i32) -> f32 { + self.assert_token_in_range(i_segment, i_token); unsafe { ffi::whisper_full_get_token_p_from_state(self.ptr, i_segment, i_token) } } } diff --git a/whisper-cpp-plus/src/stream.rs b/whisper-cpp-plus/src/stream.rs index b76d5c4..a5ef4dd 100644 --- a/whisper-cpp-plus/src/stream.rs +++ b/whisper-cpp-plus/src/stream.rs @@ -265,12 +265,10 @@ impl WhisperStream { return Ok(Vec::new()); } - // Clone params so we can set prompt_tokens pointer + // Clone params so the per-iteration prompt tokens don't leak into self.params. + // prompt_tokens() copies the tokens into the params. let mut params = self.params.clone(); - // Set prompt tokens on the clone, pointing to self.prompt_tokens. - // The prompt_tokens() method stores a raw pointer. self.prompt_tokens - // (Vec) lives on self and outlives the full() call, so this is safe. if !self.config.no_context && !self.prompt_tokens.is_empty() { params = params.prompt_tokens(&self.prompt_tokens); } diff --git a/whisper-cpp-plus/tests/real_audio.rs b/whisper-cpp-plus/tests/real_audio.rs index a42f1d3..6d9712d 100644 --- a/whisper-cpp-plus/tests/real_audio.rs +++ b/whisper-cpp-plus/tests/real_audio.rs @@ -1,5 +1,6 @@ +use std::panic::{catch_unwind, AssertUnwindSafe}; use std::path::Path; -use whisper_cpp_plus::{FullParams, SamplingStrategy, WhisperContext}; +use whisper_cpp_plus::{FullParams, SamplingStrategy, WhisperContext, WhisperState}; /// Find Whisper model (env var or default paths) fn find_whisper_model() -> Option { @@ -135,6 +136,111 @@ fn test_jfk_transcription() { } } +#[test] +fn test_jfk_segment_times_are_milliseconds() { + let Some(model_path) = find_whisper_model() else { + eprintln!( + "Skipping: model not found. Set WHISPER_TEST_MODEL_DIR or run `cargo xtask test-setup`" + ); + return; + }; + let Some(audio_path) = find_jfk_audio() else { + eprintln!("Skipping: JFK audio not found. Set WHISPER_TEST_AUDIO_DIR or run `cargo xtask test-setup`"); + return; + }; + + let audio = load_wav_file(&audio_path).expect("Failed to load JFK audio"); + let audio_ms = audio.len() as i64 * 1000 / 16000; + + let ctx = WhisperContext::new(&model_path).expect("Failed to load model"); + let result = ctx + .transcribe_with_full_params( + &audio, + FullParams::new(SamplingStrategy::Greedy { best_of: 1 }), + ) + .expect("Failed to transcribe"); + + let last = result + .segments + .last() + .expect("Should have at least one segment"); + + // jfk.wav is ~11 s. whisper.cpp reports centiseconds; if they leaked through unconverted, + // the last segment would end around 1100 instead of 11000. + assert!( + last.end_ms > audio_ms * 3 / 4 && last.end_ms <= audio_ms + 1000, + "last segment ends at {} ms, expected close to the {} ms audio length", + last.end_ms, + audio_ms + ); + assert!((last.end_seconds() - last.end_ms as f64 / 1000.0).abs() < f64::EPSILON); +} + +#[test] +fn test_state_getters_reject_out_of_range_indices() { + let Some(model_path) = find_whisper_model() else { + eprintln!( + "Skipping: model not found. Set WHISPER_TEST_MODEL_DIR or run `cargo xtask test-setup`" + ); + return; + }; + let Some(audio_path) = find_jfk_audio() else { + eprintln!("Skipping: JFK audio not found. Set WHISPER_TEST_AUDIO_DIR or run `cargo xtask test-setup`"); + return; + }; + + let audio = load_wav_file(&audio_path).expect("Failed to load JFK audio"); + let ctx = WhisperContext::new(&model_path).expect("Failed to load model"); + let mut state = WhisperState::new(&ctx).expect("Failed to create state"); + state + .full( + FullParams::new(SamplingStrategy::Greedy { best_of: 1 }), + &audio, + ) + .expect("Failed to transcribe"); + + let n_segments = state.full_n_segments(); + assert!(n_segments > 0); + let n_tokens = state.full_n_tokens(0); + assert!(n_tokens > 0); + + // In-range access works. + assert!(state.full_get_segment_text(0).is_ok()); + assert!(state.full_get_token_text(0, 0).is_ok()); + assert!(state.full_get_token_data(0, 0).is_some()); + let no_speech = state.full_get_segment_no_speech_prob(0); + assert!((0.0..=1.0).contains(&no_speech)); + + // Result/Option getters report out-of-range indices. + assert!(state.full_get_segment_text(n_segments).is_err()); + assert!(state.full_get_segment_text(-1).is_err()); + assert!(state.full_get_token_text(0, n_tokens).is_err()); + assert!(state.full_get_token_text(n_segments, 0).is_err()); + assert!(state.full_get_token_data(0, n_tokens).is_none()); + assert!(state.full_get_token_data(0, -1).is_none()); + + // Plain-value getters panic instead of reading out of bounds in C. + let panics = |f: &dyn Fn()| catch_unwind(AssertUnwindSafe(f)).is_err(); + assert!(panics(&|| { + state.full_get_segment_timestamps(n_segments); + })); + assert!(panics(&|| { + state.full_get_segment_speaker_turn_next(-1); + })); + assert!(panics(&|| { + state.full_get_segment_no_speech_prob(n_segments); + })); + assert!(panics(&|| { + state.full_n_tokens(n_segments); + })); + assert!(panics(&|| { + state.full_get_token_id(0, n_tokens); + })); + assert!(panics(&|| { + state.full_get_token_prob(0, n_tokens); + })); +} + #[test] fn test_audio_duration_handling() { let Some(model_path) = find_whisper_model() else {