1#![doc = include_str!("../README.md")]
2
3pub mod surface;
4
5#[cfg(feature = "alignment")]
6mod ctc_alignment;
7mod native_audio;
8mod native_bundles;
9mod native_device;
10#[cfg(feature = "alignment")]
11mod native_wav2vec2;
12#[cfg(feature = "alignment")]
13mod native_wav2vec2_model;
14#[cfg(feature = "candle")]
15mod native_whisper;
16#[cfg(any(feature = "silero-vad", feature = "pyannote-vad", test))]
17mod silero_vad;
18
19use std::collections::BTreeMap;
20use std::fs;
21use std::path::{Path, PathBuf};
22use std::process::{Child, Command, Output, Stdio};
23use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
24
25use serde::{Deserialize, Serialize};
26use text_transcripts::{
27 normalize_transcription_contract, TranscriptCharContract, TranscriptWordContract,
28 TranscriptionContract,
29};
30use video_analysis_core::{DetectError, Result};
31
32const CANDLE_WHISPER_AUTOREGRESSIVE_KV_CACHE_EXECUTION: &str =
33 "candle-whisper-autoregressive-kv-cache";
34const CANDLE_WHISPER_ACTIVE_ROW_TENSOR_BATCH_EXECUTION: &str =
35 "candle-whisper-active-row-tensor-batch";
36
37pub use audio_analysis_speakers::{
38 AudioRuntime, SpeakerDiarizationOptions, SpeakerDiarizationResponse, SpeakerSegmentPrediction,
39 SpeakerTranscriptAssignmentPolicy,
40};
41#[cfg(feature = "pyannote-vad")]
42pub use silero_vad::{PyannoteVadOptions, PyannoteVadTranscriptionProvider};
43#[cfg(feature = "silero-vad")]
44pub use silero_vad::{SileroVadOptions, SileroVadTranscriptionProvider};
45
46pub type SpeakerAssignmentPolicy = SpeakerTranscriptAssignmentPolicy;
48
49#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
51#[serde(rename_all = "camelCase")]
52pub struct TranscriptionPipelineRequest {
53 pub source: TranscriptionSource,
54 pub provider: TranscriptionProviderSelection,
55 #[serde(default)]
56 pub vad: VadOptions,
57 #[serde(default)]
58 pub alignment: AlignmentOptions,
59 #[serde(default)]
60 pub diarization: DiarizationOptions,
61 #[serde(default)]
62 pub output: TranscriptionOutputOptions,
63}
64
65#[derive(Debug, Clone, PartialEq)]
67pub enum TranscriptionPipelineEvent {
68 ValidationStart,
69 DecodeStart,
70 DecodeEnd {
71 duration_seconds: f64,
72 samples: usize,
73 },
74 VadStart {
75 provider: String,
76 },
77 VadEnd {
78 segments: usize,
79 windows: Option<usize>,
80 },
81 AsrStart {
82 model_id: String,
83 },
84 AsrEnd {
85 segments: usize,
86 },
87 ModelLoadStart {
88 stage: String,
89 provider: String,
90 model_id: String,
91 },
92 ModelLoadEnd {
93 stage: String,
94 provider: String,
95 model_id: String,
96 duration_seconds: f64,
97 },
98 ModelReuse {
99 stage: String,
100 provider: String,
101 model_id: String,
102 },
103 AlignmentStart {
104 model_id: String,
105 },
106 AlignmentEnd {
107 words: usize,
108 },
109 DiarizationStart {
110 provider: String,
111 },
112 DiarizationEnd {
113 speakers: usize,
114 segments: usize,
115 },
116}
117
118pub trait TranscriptionPipelineObserver {
120 fn observe(&mut self, event: TranscriptionPipelineEvent);
121
122 fn model_resolution_start(&mut self, _stage: &str, _provider: &str, _model_id: &str) {}
123
124 fn model_resolution_end(
125 &mut self,
126 _stage: &str,
127 _provider: &str,
128 _model_id: &str,
129 _source: &str,
130 ) {
131 }
132
133 fn model_download_start(&mut self, _stage: &str, _provider: &str, _model_id: &str) {}
134
135 fn model_download_end(
136 &mut self,
137 _stage: &str,
138 _provider: &str,
139 _model_id: &str,
140 _duration_seconds: f64,
141 ) {
142 }
143
144 fn cancellation_requested(&self) -> bool {
146 false
147 }
148}
149
150#[derive(Debug, Default)]
152pub struct NoopTranscriptionPipelineObserver;
153
154impl TranscriptionPipelineObserver for NoopTranscriptionPipelineObserver {
155 fn observe(&mut self, _event: TranscriptionPipelineEvent) {}
156}
157
158#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
160#[serde(rename_all = "camelCase", untagged)]
161pub enum TranscriptionSource {
162 Path {
163 path: PathBuf,
164 },
165 Samples {
166 samples: Vec<f32>,
167 #[serde(rename = "sampleRate", alias = "sample_rate")]
168 sample_rate: u32,
169 channels: u16,
170 #[serde(default)]
171 source: Option<String>,
172 },
173}
174
175impl TranscriptionSource {
176 fn path(&self) -> Result<&Path> {
177 match self {
178 Self::Path { path } => Ok(path),
179 Self::Samples { .. } => Err(invalid_request(
180 "external command transcription requires a path source",
181 )),
182 }
183 }
184}
185
186#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
188#[serde(rename_all = "camelCase", tag = "kind")]
189pub enum TranscriptionProviderSelection {
190 #[serde(rename = "candleWhisper", alias = "candle-whisper")]
191 CandleWhisper(CandleWhisperOptions),
192 #[serde(rename = "whisperCpp", alias = "whisper-cpp")]
193 WhisperCpp(WhisperCppProviderOptions),
194 #[serde(rename = "externalWhisperX", alias = "whisperx")]
195 ExternalWhisperX(WhisperXCommandOptions),
196}
197
198impl TranscriptionProviderSelection {
199 pub fn provider_id(&self) -> &'static str {
200 match self {
201 Self::CandleWhisper(_) => "candle-whisper",
202 Self::WhisperCpp(_) => "whisper-cpp",
203 Self::ExternalWhisperX(_) => "whisperx-command",
204 }
205 }
206
207 pub fn model_id(&self) -> &str {
208 match self {
209 Self::CandleWhisper(options) => &options.model_id,
210 Self::WhisperCpp(options) => &options.model_id,
211 Self::ExternalWhisperX(options) => &options.model,
212 }
213 }
214
215 pub fn task(&self) -> TranscriptionTask {
216 match self {
217 Self::CandleWhisper(options) => options.task,
218 Self::WhisperCpp(options) => options.task,
219 Self::ExternalWhisperX(options) => options.task,
220 }
221 }
222}
223
224#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
226#[serde(rename_all = "lowercase")]
227pub enum TranscriptionTask {
228 #[default]
229 Transcribe,
230 Translate,
231}
232
233impl TranscriptionTask {
234 pub fn as_whisper_task(self) -> &'static str {
235 match self {
236 Self::Transcribe => "transcribe",
237 Self::Translate => "translate",
238 }
239 }
240
241 pub fn output_language_hint(self) -> Option<&'static str> {
242 match self {
243 Self::Transcribe => None,
244 Self::Translate => Some("en"),
245 }
246 }
247}
248
249#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
251#[serde(rename_all = "camelCase")]
252pub struct CandleWhisperOptions {
253 #[serde(default = "default_candle_whisper_model")]
254 pub model_id: String,
255 #[serde(default)]
256 pub task: TranscriptionTask,
257 #[serde(default)]
258 pub language: Option<String>,
259 #[serde(default)]
260 pub device: NativeDevicePreference,
261 #[serde(default)]
262 pub compute_type: CandleWhisperComputeType,
263 #[serde(default)]
264 pub model_bundle: Option<PathBuf>,
265 #[serde(default)]
266 pub model_dir: Option<PathBuf>,
267 #[serde(default)]
268 pub model_cache_only: bool,
269 #[serde(default)]
270 pub batch_chunks: bool,
271 #[serde(default)]
272 pub max_batch_size: Option<usize>,
273 #[serde(default)]
274 pub decode_runtime: CandleWhisperDecodeRuntime,
275}
276
277impl Default for CandleWhisperOptions {
278 fn default() -> Self {
279 Self {
280 model_id: default_candle_whisper_model(),
281 task: TranscriptionTask::Transcribe,
282 language: None,
283 device: NativeDevicePreference::Auto,
284 compute_type: CandleWhisperComputeType::Automatic,
285 model_bundle: None,
286 model_dir: None,
287 model_cache_only: false,
288 batch_chunks: true,
289 max_batch_size: Some(4),
290 decode_runtime: CandleWhisperDecodeRuntime::AutoregressiveKvCache,
291 }
292 }
293}
294
295fn default_candle_whisper_model() -> String {
296 "openai/whisper-large-v3-turbo".to_string()
297}
298
299#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
301#[serde(rename_all = "lowercase")]
302pub enum CandleWhisperComputeType {
303 #[default]
304 #[serde(alias = "auto")]
305 Automatic,
306 #[serde(alias = "float16")]
307 Fp16,
308 #[serde(alias = "float32")]
309 Fp32,
310}
311
312impl CandleWhisperComputeType {
313 pub fn as_str(self) -> &'static str {
314 match self {
315 Self::Automatic => "automatic",
316 Self::Fp16 => "fp16",
317 Self::Fp32 => "fp32",
318 }
319 }
320
321 pub(crate) fn resolve_for_device(self, cuda_active: bool) -> Result<Self> {
322 match (self, cuda_active) {
323 (Self::Automatic, true) => Ok(Self::Fp16),
324 (Self::Automatic, false) => Ok(Self::Fp32),
325 (Self::Fp16, true) => Ok(Self::Fp16),
326 (Self::Fp16, false) => Err(setup_error(
327 "native Candle Whisper compute type fp16 requires a CUDA device; use automatic or fp32 for CPU execution",
328 )),
329 (Self::Fp32, _) => Ok(Self::Fp32),
330 }
331 }
332
333 #[allow(dead_code)]
334 pub(crate) fn setup_fallback_eligible(self) -> bool {
335 matches!(self, Self::Automatic)
336 }
337}
338
339#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
341#[serde(rename_all = "camelCase")]
342pub enum CandleWhisperDecodeRuntime {
343 #[default]
345 AutoregressiveKvCache,
346 ActiveRowTensorBatch,
348}
349
350impl CandleWhisperDecodeRuntime {
351 pub fn execution_id(self) -> &'static str {
352 match self {
353 Self::AutoregressiveKvCache => CANDLE_WHISPER_AUTOREGRESSIVE_KV_CACHE_EXECUTION,
354 Self::ActiveRowTensorBatch => CANDLE_WHISPER_ACTIVE_ROW_TENSOR_BATCH_EXECUTION,
355 }
356 }
357
358 pub fn is_supported(self) -> bool {
359 matches!(
360 self,
361 Self::AutoregressiveKvCache | Self::ActiveRowTensorBatch
362 )
363 }
364}
365
366#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
368#[serde(rename_all = "camelCase")]
369pub struct WhisperCppProviderOptions {
370 #[serde(default = "default_whisper_cpp_model")]
371 pub model_id: String,
372 #[serde(default)]
373 pub task: TranscriptionTask,
374 #[serde(default)]
375 pub language: Option<String>,
376 #[serde(default)]
377 pub model_path: Option<PathBuf>,
378}
379
380impl Default for WhisperCppProviderOptions {
381 fn default() -> Self {
382 Self {
383 model_id: default_whisper_cpp_model(),
384 task: TranscriptionTask::Transcribe,
385 language: None,
386 model_path: None,
387 }
388 }
389}
390
391fn default_whisper_cpp_model() -> String {
392 "large-v3-turbo".to_string()
393}
394
395#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
397#[serde(rename_all = "lowercase")]
398pub enum NativeDevicePreference {
399 #[default]
400 Auto,
401 Cpu,
402 Cuda,
403}
404
405#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
407#[serde(rename_all = "camelCase")]
408pub struct WhisperXCommandOptions {
409 pub command: PathBuf,
410 pub model: String,
411 #[serde(default)]
412 pub task: TranscriptionTask,
413 #[serde(default)]
414 pub language: Option<String>,
415 pub device: WhisperXDevice,
416 #[serde(default)]
417 pub compute_type: Option<String>,
418 #[serde(default)]
419 pub batch_size: Option<usize>,
420 #[serde(default)]
421 pub align_model: Option<String>,
422 #[serde(default)]
423 pub model_dir: Option<PathBuf>,
424 #[serde(default)]
425 pub model_cache_only: bool,
426 #[serde(default)]
427 pub no_align: bool,
428 #[serde(default)]
429 pub interpolate_method: AlignmentInterpolationMethod,
430 #[serde(default)]
431 pub return_char_alignments: bool,
432 #[serde(default)]
433 pub diarize: bool,
434 #[serde(default)]
435 pub min_speakers: Option<usize>,
436 #[serde(default)]
437 pub max_speakers: Option<usize>,
438 #[serde(default)]
439 pub hf_token_env: Option<String>,
440 #[serde(default)]
441 pub output_dir: Option<PathBuf>,
442 #[serde(default)]
443 pub timeout_seconds: Option<u64>,
444 #[serde(default)]
445 pub extra_args: Vec<String>,
446}
447
448pub type WhisperXOptions = WhisperXCommandOptions;
450
451impl Default for WhisperXCommandOptions {
452 fn default() -> Self {
453 Self {
454 command: PathBuf::from("whisperx"),
455 model: "large-v2".to_string(),
456 task: TranscriptionTask::Transcribe,
457 language: None,
458 device: WhisperXDevice::Cpu,
459 compute_type: None,
460 batch_size: None,
461 align_model: None,
462 model_dir: None,
463 model_cache_only: false,
464 no_align: false,
465 interpolate_method: AlignmentInterpolationMethod::Nearest,
466 return_char_alignments: false,
467 diarize: false,
468 min_speakers: None,
469 max_speakers: None,
470 hf_token_env: None,
471 output_dir: None,
472 timeout_seconds: None,
473 extra_args: Vec::new(),
474 }
475 }
476}
477
478#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
480#[serde(rename_all = "lowercase")]
481pub enum WhisperXDevice {
482 Cpu,
483 Cuda,
484}
485
486impl WhisperXDevice {
487 fn as_str(self) -> &'static str {
488 match self {
489 Self::Cpu => "cpu",
490 Self::Cuda => "cuda",
491 }
492 }
493}
494
495#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
497#[serde(rename_all = "lowercase")]
498pub enum AlignmentInterpolationMethod {
499 #[default]
500 Nearest,
501 Linear,
502 Ignore,
503}
504
505impl AlignmentInterpolationMethod {
506 pub fn as_whisperx_arg(self) -> &'static str {
507 match self {
508 Self::Nearest => "nearest",
509 Self::Linear => "linear",
510 Self::Ignore => "ignore",
511 }
512 }
513}
514
515#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
517#[serde(rename_all = "camelCase")]
518pub struct VadOptions {
519 #[serde(default = "default_vad_enabled")]
520 pub enabled: bool,
521 #[serde(default = "default_vad_threshold")]
522 pub rms_threshold: f32,
523 #[serde(default = "default_vad_frame_seconds")]
524 pub frame_seconds: f64,
525 #[serde(default = "default_vad_hop_seconds")]
526 pub hop_seconds: f64,
527 #[serde(default = "default_vad_min_speech_seconds")]
528 pub min_speech_seconds: f64,
529 #[serde(default = "default_vad_padding_seconds")]
530 pub padding_seconds: f64,
531 #[serde(default = "default_vad_merge_gap_seconds")]
532 pub merge_gap_seconds: f64,
533 #[serde(default = "default_vad_max_chunk_seconds")]
534 pub max_chunk_seconds: f64,
535}
536
537impl Default for VadOptions {
538 fn default() -> Self {
539 Self {
540 enabled: true,
541 rms_threshold: default_vad_threshold(),
542 frame_seconds: default_vad_frame_seconds(),
543 hop_seconds: default_vad_hop_seconds(),
544 min_speech_seconds: default_vad_min_speech_seconds(),
545 padding_seconds: default_vad_padding_seconds(),
546 merge_gap_seconds: default_vad_merge_gap_seconds(),
547 max_chunk_seconds: default_vad_max_chunk_seconds(),
548 }
549 }
550}
551
552fn default_vad_enabled() -> bool {
553 true
554}
555fn default_vad_threshold() -> f32 {
556 0.01
557}
558fn default_vad_frame_seconds() -> f64 {
559 0.03
560}
561fn default_vad_hop_seconds() -> f64 {
562 0.01
563}
564fn default_vad_min_speech_seconds() -> f64 {
565 0.08
566}
567fn default_vad_padding_seconds() -> f64 {
568 0.02
569}
570fn default_vad_merge_gap_seconds() -> f64 {
571 0.05
572}
573fn default_vad_max_chunk_seconds() -> f64 {
574 30.0
575}
576
577#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
579#[serde(rename_all = "camelCase")]
580pub struct AlignmentOptions {
581 #[serde(default)]
582 pub enabled: bool,
583 #[serde(default = "default_alignment_model")]
584 pub model_id: String,
585 #[serde(default = "default_alignment_device")]
586 pub device: NativeDevicePreference,
587 #[serde(default)]
588 pub model_bundle: Option<PathBuf>,
589 #[serde(default)]
590 pub model_dir: Option<PathBuf>,
591 #[serde(default)]
592 pub model_cache_only: bool,
593 #[serde(default)]
594 pub interpolate_method: AlignmentInterpolationMethod,
595 #[serde(default)]
596 pub return_char_alignments: bool,
597}
598
599impl Default for AlignmentOptions {
600 fn default() -> Self {
601 Self {
602 enabled: false,
603 model_id: default_alignment_model(),
604 device: default_alignment_device(),
605 model_bundle: None,
606 model_dir: None,
607 model_cache_only: false,
608 interpolate_method: AlignmentInterpolationMethod::Nearest,
609 return_char_alignments: false,
610 }
611 }
612}
613
614fn default_alignment_model() -> String {
615 "facebook/wav2vec2-base-960h".to_string()
616}
617
618fn default_alignment_device() -> NativeDevicePreference {
619 NativeDevicePreference::Cpu
620}
621
622#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
624#[serde(rename_all = "camelCase")]
625pub struct DiarizationOptions {
626 #[serde(default)]
627 pub enabled: bool,
628 #[serde(default, flatten)]
629 pub speaker: SpeakerDiarizationOptions,
630}
631
632impl std::ops::Deref for DiarizationOptions {
633 type Target = SpeakerDiarizationOptions;
634
635 fn deref(&self) -> &Self::Target {
636 &self.speaker
637 }
638}
639
640impl std::ops::DerefMut for DiarizationOptions {
641 fn deref_mut(&mut self) -> &mut Self::Target {
642 &mut self.speaker
643 }
644}
645
646#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
648#[serde(rename_all = "camelCase")]
649pub struct TranscriptionOutputOptions {
650 #[serde(default = "default_output_formats")]
651 pub formats: Vec<String>,
652}
653
654impl Default for TranscriptionOutputOptions {
655 fn default() -> Self {
656 Self {
657 formats: default_output_formats(),
658 }
659 }
660}
661
662fn default_output_formats() -> Vec<String> {
663 vec!["json".to_string()]
664}
665
666#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
668#[serde(rename_all = "camelCase")]
669pub struct TranscriptionArtifact {
670 pub kind: String,
671 pub path: PathBuf,
672}
673
674#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
676#[serde(rename_all = "camelCase")]
677pub struct TranscriptionPipelineResponse {
678 pub accepted: bool,
679 pub operation: String,
680 pub provider: String,
681 pub model_id: String,
682 pub transcript: TranscriptionContract,
683 pub vad_segments: Vec<SpeechActivitySegment>,
684 pub alignment: Option<AlignmentSummary>,
685 pub diarization: Option<SpeakerDiarizationResponse>,
686 pub artifacts: Vec<TranscriptionArtifact>,
687 pub diagnostics: Vec<String>,
688}
689
690#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
692#[serde(rename_all = "camelCase")]
693pub struct TranscriptionProviderPlan {
694 pub provider_id: String,
695 pub external_runtime: bool,
696 pub wasm_supported: bool,
697 pub primary: bool,
698 pub setup: Vec<String>,
699 pub diagnostics: Vec<String>,
700}
701
702#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
704#[serde(rename_all = "camelCase")]
705pub struct LoadedAudio {
706 pub samples: Vec<f32>,
707 pub sample_rate: u32,
708 pub channels: u16,
709 #[serde(default)]
710 pub source: Option<String>,
711}
712
713impl LoadedAudio {
714 pub fn mono_16khz_from_source(source: &TranscriptionSource) -> Result<Self> {
715 native_audio::mono_16khz_from_source(source)
716 }
717
718 pub fn duration_seconds(&self) -> f64 {
719 if self.sample_rate == 0 || self.channels == 0 {
720 return 0.0;
721 }
722 self.samples.len() as f64 / self.channels as f64 / self.sample_rate as f64
723 }
724}
725
726#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
728#[serde(rename_all = "camelCase")]
729pub struct SpeechActivitySegment {
730 pub start_seconds: f64,
731 pub end_seconds: f64,
732 pub score: f32,
733}
734
735impl SpeechActivitySegment {
736 pub fn new(start_seconds: f64, end_seconds: f64, score: f32) -> Result<Self> {
737 let segment = Self {
738 start_seconds,
739 end_seconds,
740 score,
741 };
742 segment.validate()?;
743 Ok(segment)
744 }
745
746 pub fn validate(&self) -> Result<()> {
747 if !self.start_seconds.is_finite()
748 || !self.end_seconds.is_finite()
749 || !self.score.is_finite()
750 {
751 return Err(invalid_request(
752 "speech activity segment values must be finite",
753 ));
754 }
755 if self.start_seconds < 0.0 || self.end_seconds <= self.start_seconds {
756 return Err(invalid_request(
757 "speech activity segment must have non-negative start and positive duration",
758 ));
759 }
760 Ok(())
761 }
762}
763
764#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
766#[serde(rename_all = "camelCase")]
767pub struct AsrRequest {
768 pub audio: LoadedAudio,
769 pub chunks: Vec<SpeechActivitySegment>,
770 #[serde(default)]
771 pub task: TranscriptionTask,
772 pub language: Option<String>,
773 pub model_id: String,
774}
775
776#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
778#[serde(rename_all = "camelCase")]
779pub struct AsrResponse {
780 pub model_id: String,
781 pub language: Option<String>,
782 pub transcript: TranscriptionContract,
783 #[serde(default)]
784 pub diagnostics: Vec<String>,
785}
786
787#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
789#[serde(rename_all = "camelCase")]
790pub struct AlignmentRequest {
791 pub audio: LoadedAudio,
792 pub transcript: TranscriptionContract,
793 pub language: Option<String>,
794 pub model_id: String,
795}
796
797#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
799#[serde(rename_all = "camelCase")]
800pub struct AlignmentResponse {
801 pub model_id: String,
802 pub words: Vec<AlignedWord>,
803 #[serde(default)]
804 pub chars: Vec<AlignedChar>,
805 #[serde(default)]
806 pub diagnostics: Vec<String>,
807}
808
809#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
811#[serde(rename_all = "camelCase")]
812pub struct AlignedWord {
813 pub segment_index: u64,
814 pub word_index: usize,
815 pub text: String,
816 pub start_seconds: f64,
817 pub end_seconds: f64,
818 #[serde(default)]
819 pub confidence: Option<f32>,
820}
821
822#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
824#[serde(rename_all = "camelCase")]
825pub struct AlignedChar {
826 pub segment_index: u64,
827 pub char_index: usize,
828 pub character: String,
829 #[serde(default)]
830 pub start_seconds: Option<f64>,
831 #[serde(default)]
832 pub end_seconds: Option<f64>,
833 #[serde(default)]
834 pub confidence: Option<f32>,
835}
836
837#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
839#[serde(rename_all = "camelCase")]
840pub struct AlignmentSummary {
841 pub provider: String,
842 pub model_id: String,
843 pub word_count: usize,
844}
845
846#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
848#[serde(rename_all = "camelCase")]
849pub struct VadRequest {
850 pub audio: LoadedAudio,
851 pub options: VadOptions,
852}
853
854#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
856#[serde(rename_all = "camelCase")]
857pub struct VadResponse {
858 pub segments: Vec<SpeechActivitySegment>,
859 #[serde(default)]
860 pub diagnostics: Vec<String>,
861}
862
863pub trait AudioTranscriptionProvider {
865 fn provider_id(&self) -> &str;
866 fn transcribe(&mut self, request: AsrRequest) -> Result<AsrResponse>;
867
868 fn transcribe_with_observer(
869 &mut self,
870 request: AsrRequest,
871 _observer: &mut dyn TranscriptionPipelineObserver,
872 ) -> Result<AsrResponse> {
873 self.transcribe(request)
874 }
875}
876
877pub trait ForcedAlignmentProvider {
879 fn provider_id(&self) -> &str;
880 fn align(&mut self, request: AlignmentRequest) -> Result<AlignmentResponse>;
881
882 fn align_with_observer(
883 &mut self,
884 request: AlignmentRequest,
885 _observer: &mut dyn TranscriptionPipelineObserver,
886 ) -> Result<AlignmentResponse> {
887 self.align(request)
888 }
889}
890
891pub trait TranscriptionVadProvider {
893 fn provider_id(&self) -> &str;
894 fn detect_speech(&mut self, request: VadRequest) -> Result<VadResponse>;
895}
896
897pub trait TranscriptDiarizationProvider {
899 fn provider_id(&self) -> &str;
900 fn diarize(
901 &mut self,
902 audio: LoadedAudio,
903 transcript: &TranscriptionContract,
904 options: &DiarizationOptions,
905 ) -> Result<SpeakerDiarizationResponse>;
906}
907
908#[cfg(feature = "diarization")]
910#[derive(Debug, Clone, Default)]
911pub struct NativeSpeakerDiarizationProvider;
912
913#[cfg(feature = "diarization")]
914impl TranscriptDiarizationProvider for NativeSpeakerDiarizationProvider {
915 fn provider_id(&self) -> &str {
916 "native-speaker-diarization"
917 }
918
919 fn diarize(
920 &mut self,
921 audio: LoadedAudio,
922 transcript: &TranscriptionContract,
923 options: &DiarizationOptions,
924 ) -> Result<SpeakerDiarizationResponse> {
925 native_audio::validate_loaded_audio(&audio)?;
926 if options.is_pyannote_model() {
927 #[cfg(feature = "pyannote-diarization")]
928 {
929 let mut provider = PyannoteCommunityTranscriptDiarizationProvider;
930 return provider.diarize(audio, transcript, options);
931 }
932 #[cfg(not(feature = "pyannote-diarization"))]
933 {
934 return Err(setup_error(
935 "native pyannote diarization requires the pyannote-diarization feature",
936 ));
937 }
938 }
939 if options.speaker_embedding_model_bundle.is_some() {
940 return diarize_with_onnx_speaker_embeddings(audio, transcript, options);
941 }
942 let spans = speech_spans_from_transcript(transcript, audio.duration_seconds())?;
943 if !spans.is_empty() {
944 let speaker_audio =
945 audio_analysis_speakers::SpeakerAudio::mono(&audio.samples, audio.sample_rate)?;
946 let embedder = audio_analysis_speakers::SpectralSpeakerEmbedder::default();
947 let vad = TranscriptSpeechSpanVad { spans };
948 let mut diarizer = audio_analysis_speakers::WindowedSpeakerDiarizer::new(embedder, vad)
949 .cluster_threshold(0.95)?
950 .speaker_bounds(options.min_speakers, options.max_speakers)?;
951 let result =
952 audio_analysis_speakers::SpeakerDiarizer::diarize(&mut diarizer, &speaker_audio)?;
953 return Ok(SpeakerDiarizationResponse {
954 accepted: true,
955 operation: "audio.speakers.diarize".to_string(),
956 model_id: options.model_id.clone(),
957 runtime: audio_analysis_speakers::AudioRuntime::Heuristic,
958 segments: stable_speaker_predictions_from_diarization(result.segments)?,
959 speaker_embeddings: None,
960 diagnostics: Vec::new(),
961 });
962 }
963
964 let speaker_audio =
965 audio_analysis_speakers::SpeakerAudio::mono(&audio.samples, audio.sample_rate)?;
966 let embedder = audio_analysis_speakers::SpectralSpeakerEmbedder::default();
967 let vad_config = audio_analysis_speakers::EnergyVadConfig::default();
968 let vad = audio_analysis_speakers::EnergyVoiceActivityDetector::new(vad_config)?;
969 let mut diarizer = audio_analysis_speakers::WindowedSpeakerDiarizer::new(embedder, vad)
970 .cluster_threshold(0.95)?
971 .speaker_bounds(options.min_speakers, options.max_speakers)?;
972 let result =
973 audio_analysis_speakers::SpeakerDiarizer::diarize(&mut diarizer, &speaker_audio)?;
974 Ok(SpeakerDiarizationResponse {
975 accepted: true,
976 operation: "audio.speakers.diarize".to_string(),
977 model_id: options.model_id.clone(),
978 runtime: audio_analysis_speakers::AudioRuntime::Heuristic,
979 segments: stable_speaker_predictions_from_diarization(result.segments)?,
980 speaker_embeddings: None,
981 diagnostics: Vec::new(),
982 })
983 }
984}
985
986#[cfg(all(feature = "diarization", feature = "pyannote-diarization"))]
988#[derive(Debug, Clone, Default)]
989pub struct PyannoteCommunityTranscriptDiarizationProvider;
990
991#[cfg(all(feature = "diarization", feature = "pyannote-diarization"))]
992impl TranscriptDiarizationProvider for PyannoteCommunityTranscriptDiarizationProvider {
993 fn provider_id(&self) -> &str {
994 "pyannote-community-diarization"
995 }
996
997 fn diarize(
998 &mut self,
999 audio: LoadedAudio,
1000 _transcript: &TranscriptionContract,
1001 options: &DiarizationOptions,
1002 ) -> Result<SpeakerDiarizationResponse> {
1003 native_audio::validate_loaded_audio(&audio)?;
1004 if !options.is_pyannote_model() {
1005 return Err(invalid_request(format!(
1006 "pyannote community diarization provider does not support model `{}`",
1007 options.model_id
1008 )));
1009 }
1010 let bundle_path = options.pyannote_model_bundle.clone().ok_or_else(|| {
1011 setup_error(
1012 "native pyannote diarization requires --diarization-model-bundle or DiarizationOptions.pyannote_model_bundle",
1013 )
1014 })?;
1015 let speaker_audio =
1016 audio_analysis_speakers::SpeakerAudio::mono(&audio.samples, audio.sample_rate)?;
1017 let mut diarizer = audio_analysis_speakers::PyannoteCommunityDiarizer::from_config(
1018 audio_analysis_speakers::PyannoteCommunityDiarizationConfig {
1019 bundle_path,
1020 manifest_file: options.pyannote_manifest_file.clone(),
1021 segmentation_model_file: options.pyannote_segmentation_model_file.clone(),
1022 embedding_model_file: options.pyannote_embedding_model_file.clone(),
1023 plda_transform_file: options.pyannote_plda_transform_file.clone(),
1024 plda_model_file: options.pyannote_plda_model_file.clone(),
1025 clustering_config_file: options.pyannote_clustering_config_file.clone(),
1026 min_speakers: options.min_speakers,
1027 max_speakers: options.max_speakers,
1028 return_speaker_embeddings: options.return_speaker_embeddings,
1029 },
1030 )?;
1031 let mut result = diarizer.diarize(&speaker_audio)?.response;
1032 result.model_id = options.model_id.clone();
1033 Ok(result)
1034 }
1035}
1036
1037#[cfg(feature = "diarization")]
1038fn diarize_with_onnx_speaker_embeddings(
1039 audio: LoadedAudio,
1040 transcript: &TranscriptionContract,
1041 options: &DiarizationOptions,
1042) -> Result<SpeakerDiarizationResponse> {
1043 let config = options.onnx_speaker_embedding_config()?;
1044 let speaker_audio =
1045 audio_analysis_speakers::SpeakerAudio::mono(&audio.samples, audio.sample_rate)?;
1046 let embedder = audio_analysis_speakers::OnnxSpeakerEmbedder::from_config(config)?;
1047 let spans = speech_spans_from_transcript(transcript, audio.duration_seconds())?;
1048 let result = if spans.is_empty() {
1049 let vad = audio_analysis_speakers::EnergyVoiceActivityDetector::default();
1050 let mut diarizer = audio_analysis_speakers::WindowedSpeakerDiarizer::new(embedder, vad)
1051 .cluster_threshold(0.95)?
1052 .speaker_bounds(options.min_speakers, options.max_speakers)?;
1053 audio_analysis_speakers::SpeakerDiarizer::diarize(&mut diarizer, &speaker_audio)?
1054 } else {
1055 let vad = TranscriptSpeechSpanVad { spans };
1056 let mut diarizer = audio_analysis_speakers::WindowedSpeakerDiarizer::new(embedder, vad)
1057 .cluster_threshold(0.95)?
1058 .speaker_bounds(options.min_speakers, options.max_speakers)?;
1059 audio_analysis_speakers::SpeakerDiarizer::diarize(&mut diarizer, &speaker_audio)?
1060 };
1061 Ok(SpeakerDiarizationResponse {
1062 accepted: true,
1063 operation: "audio.speakers.diarize".to_string(),
1064 model_id: options.model_id.clone(),
1065 runtime: audio_analysis_speakers::AudioRuntime::Onnx,
1066 segments: stable_speaker_predictions_from_diarization(result.segments)?,
1067 speaker_embeddings: None,
1068 diagnostics: Vec::new(),
1069 })
1070}
1071
1072#[derive(Debug, Clone, Default)]
1074pub struct EnergyVadTranscriptionProvider;
1075
1076impl TranscriptionVadProvider for EnergyVadTranscriptionProvider {
1077 fn provider_id(&self) -> &str {
1078 "energy-vad"
1079 }
1080
1081 fn detect_speech(&mut self, request: VadRequest) -> Result<VadResponse> {
1082 let segments = energy_vad_segments(&request.audio, &request.options)?;
1083 Ok(VadResponse {
1084 segments,
1085 diagnostics: vec!["deterministic energy VAD completed".to_string()],
1086 })
1087 }
1088}
1089
1090#[derive(Debug, Clone, Default)]
1092pub struct CandleWhisperTranscriber {
1093 pub options: CandleWhisperOptions,
1094}
1095
1096impl CandleWhisperTranscriber {
1097 pub fn new(options: CandleWhisperOptions) -> Self {
1098 Self { options }
1099 }
1100}
1101
1102impl AudioTranscriptionProvider for CandleWhisperTranscriber {
1103 fn provider_id(&self) -> &str {
1104 "candle-whisper"
1105 }
1106
1107 fn transcribe(&mut self, request: AsrRequest) -> Result<AsrResponse> {
1108 let mut observer = NoopTranscriptionPipelineObserver;
1109 self.transcribe_with_observer(request, &mut observer)
1110 }
1111
1112 fn transcribe_with_observer(
1113 &mut self,
1114 request: AsrRequest,
1115 observer: &mut dyn TranscriptionPipelineObserver,
1116 ) -> Result<AsrResponse> {
1117 validate_asr_request(&request)?;
1118 validate_candle_setup(&self.options)?;
1119 let _ = &observer;
1120 #[cfg(feature = "candle")]
1121 {
1122 let chunk_count = request.chunks.len();
1123 let model_id = request.model_id.clone();
1124 let mut response =
1125 native_whisper::transcribe_with_load_observer(&self.options, request, |event| {
1126 match event {
1127 native_whisper::WhisperModelResolutionEvent::ResolutionStart => {
1128 observer.model_resolution_start("asr", "candle-whisper", &model_id);
1129 }
1130 native_whisper::WhisperModelResolutionEvent::ResolutionEnd { source } => {
1131 observer.model_resolution_end(
1132 "asr",
1133 "candle-whisper",
1134 &model_id,
1135 source,
1136 );
1137 }
1138 native_whisper::WhisperModelResolutionEvent::DownloadStart => {
1139 observer.model_download_start("asr", "hugging-face", &model_id);
1140 }
1141 native_whisper::WhisperModelResolutionEvent::DownloadEnd {
1142 duration_seconds,
1143 } => observer.model_download_end(
1144 "asr",
1145 "hugging-face",
1146 &model_id,
1147 duration_seconds,
1148 ),
1149 native_whisper::WhisperModelResolutionEvent::LoadStart => {
1150 observer.observe(TranscriptionPipelineEvent::ModelLoadStart {
1151 stage: "asr".to_string(),
1152 provider: "candle-whisper".to_string(),
1153 model_id: model_id.clone(),
1154 });
1155 }
1156 native_whisper::WhisperModelResolutionEvent::LoadEnd {
1157 duration_seconds,
1158 } => observer.observe(TranscriptionPipelineEvent::ModelLoadEnd {
1159 stage: "asr".to_string(),
1160 provider: "candle-whisper".to_string(),
1161 model_id: model_id.clone(),
1162 duration_seconds,
1163 }),
1164 }
1165 ensure_pipeline_active(observer)
1166 })?;
1167 extend_missing_candle_batch_diagnostics(
1168 &mut response.diagnostics,
1169 &self.options,
1170 chunk_count,
1171 );
1172 Ok(response)
1173 }
1174 #[cfg(not(feature = "candle"))]
1175 {
1176 Err(unsupported_runtime(format!(
1177 "Candle Whisper requested for `{}` but the binary lacks the `candle` feature; {}; build with `candle` for native execution and `model-bundles` for Hugging Face cache resolution",
1178 request.model_id,
1179 candle_whisper_setup_context(&self.options)
1180 )))
1181 }
1182 }
1183}
1184
1185#[derive(Default)]
1193pub struct ReusableCandleWhisperTranscriber {
1194 pub options: CandleWhisperOptions,
1195 #[cfg(feature = "candle")]
1196 session: Option<native_whisper::ReusableCandleWhisperSession>,
1197}
1198
1199impl ReusableCandleWhisperTranscriber {
1200 pub fn new(options: CandleWhisperOptions) -> Self {
1201 Self {
1202 options,
1203 #[cfg(feature = "candle")]
1204 session: None,
1205 }
1206 }
1207}
1208
1209impl std::fmt::Debug for ReusableCandleWhisperTranscriber {
1210 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1211 formatter
1212 .debug_struct("ReusableCandleWhisperTranscriber")
1213 .field("options", &self.options)
1214 .finish_non_exhaustive()
1215 }
1216}
1217
1218impl AudioTranscriptionProvider for ReusableCandleWhisperTranscriber {
1219 fn provider_id(&self) -> &str {
1220 "candle-whisper"
1221 }
1222
1223 fn transcribe(&mut self, request: AsrRequest) -> Result<AsrResponse> {
1224 let mut observer = NoopTranscriptionPipelineObserver;
1225 self.transcribe_with_observer(request, &mut observer)
1226 }
1227
1228 fn transcribe_with_observer(
1229 &mut self,
1230 request: AsrRequest,
1231 observer: &mut dyn TranscriptionPipelineObserver,
1232 ) -> Result<AsrResponse> {
1233 validate_asr_request(&request)?;
1234 validate_candle_setup(&self.options)?;
1235 let _ = &observer;
1236 #[cfg(feature = "candle")]
1237 {
1238 let chunk_count = request.chunks.len();
1239 let model_id = request.model_id.clone();
1240 let mut response = native_whisper::ReusableCandleWhisperSession::transcribe(
1241 &mut self.session,
1242 &self.options,
1243 request,
1244 |event| {
1245 match event {
1246 native_whisper::ReusableCandleWhisperSessionEvent::ResolutionStart => {
1247 observer.model_resolution_start("asr", "candle-whisper", &model_id);
1248 }
1249 native_whisper::ReusableCandleWhisperSessionEvent::ResolutionEnd {
1250 source,
1251 } => {
1252 observer.model_resolution_end(
1253 "asr",
1254 "candle-whisper",
1255 &model_id,
1256 source,
1257 );
1258 }
1259 native_whisper::ReusableCandleWhisperSessionEvent::DownloadStart => {
1260 observer.model_download_start("asr", "hugging-face", &model_id);
1261 }
1262 native_whisper::ReusableCandleWhisperSessionEvent::DownloadEnd {
1263 duration_seconds,
1264 } => {
1265 observer.model_download_end(
1266 "asr",
1267 "hugging-face",
1268 &model_id,
1269 duration_seconds,
1270 );
1271 }
1272 native_whisper::ReusableCandleWhisperSessionEvent::LoadStart => {
1273 observer.observe(TranscriptionPipelineEvent::ModelLoadStart {
1274 stage: "asr".to_string(),
1275 provider: "candle-whisper".to_string(),
1276 model_id: model_id.clone(),
1277 });
1278 }
1279 native_whisper::ReusableCandleWhisperSessionEvent::LoadEnd {
1280 duration_seconds,
1281 } => {
1282 observer.observe(TranscriptionPipelineEvent::ModelLoadEnd {
1283 stage: "asr".to_string(),
1284 provider: "candle-whisper".to_string(),
1285 model_id: model_id.clone(),
1286 duration_seconds,
1287 });
1288 }
1289 native_whisper::ReusableCandleWhisperSessionEvent::Reuse => {
1290 observer.observe(TranscriptionPipelineEvent::ModelReuse {
1291 stage: "asr".to_string(),
1292 provider: "candle-whisper".to_string(),
1293 model_id: model_id.clone(),
1294 });
1295 }
1296 }
1297 ensure_pipeline_active(observer)
1298 },
1299 )?;
1300 extend_missing_candle_batch_diagnostics(
1301 &mut response.diagnostics,
1302 &self.options,
1303 chunk_count,
1304 );
1305 Ok(response)
1306 }
1307 #[cfg(not(feature = "candle"))]
1308 {
1309 Err(unsupported_runtime(format!(
1310 "Candle Whisper requested for `{}` but the binary lacks the `candle` feature; {}; build with `candle` for native execution and `model-bundles` for Hugging Face cache resolution",
1311 request.model_id,
1312 candle_whisper_setup_context(&self.options)
1313 )))
1314 }
1315 }
1316}
1317
1318fn candle_whisper_setup_context(options: &CandleWhisperOptions) -> String {
1319 let model_location = options
1320 .model_bundle
1321 .as_ref()
1322 .map(|path| format!("--whisper-bundle={}", path.display()))
1323 .or_else(|| {
1324 options
1325 .model_dir
1326 .as_ref()
1327 .map(|path| format!("--model-dir={}", path.display()))
1328 })
1329 .unwrap_or_else(|| "--model-dir=<default huggingface cache>".to_string());
1330 format!("{model_location}; cache-only={}", options.model_cache_only)
1331}
1332
1333#[derive(Debug, Clone, Default)]
1335pub struct WhisperCppTranscriber {
1336 pub options: WhisperCppProviderOptions,
1337}
1338
1339impl AudioTranscriptionProvider for WhisperCppTranscriber {
1340 fn provider_id(&self) -> &str {
1341 "whisper-cpp"
1342 }
1343
1344 fn transcribe(&mut self, _request: AsrRequest) -> Result<AsrResponse> {
1345 let Some(model_path) = &self.options.model_path else {
1346 return Err(setup_error("required whisper.cpp model path is missing"));
1347 };
1348 if !model_path.exists() {
1349 return Err(setup_error(format!(
1350 "required whisper.cpp model `{}` is missing",
1351 model_path.display()
1352 )));
1353 }
1354 Err(unsupported_runtime(
1355 "whisper.cpp compatibility provider is not the primary transcription path",
1356 ))
1357 }
1358}
1359
1360#[derive(Debug, Clone, Default)]
1362pub struct CtcForcedAligner {
1363 pub options: AlignmentOptions,
1364}
1365
1366impl ForcedAlignmentProvider for CtcForcedAligner {
1367 fn provider_id(&self) -> &str {
1368 "ctc-forced-aligner"
1369 }
1370
1371 fn align(&mut self, request: AlignmentRequest) -> Result<AlignmentResponse> {
1372 let mut observer = NoopTranscriptionPipelineObserver;
1373 self.align_with_observer(request, &mut observer)
1374 }
1375
1376 fn align_with_observer(
1377 &mut self,
1378 request: AlignmentRequest,
1379 observer: &mut dyn TranscriptionPipelineObserver,
1380 ) -> Result<AlignmentResponse> {
1381 let _ = &observer;
1382 #[cfg(feature = "alignment")]
1383 {
1384 ctc_alignment::align_with_observer(&self.options, request, observer)
1385 }
1386 #[cfg(not(feature = "alignment"))]
1387 {
1388 validate_alignment_setup(&self.options)?;
1389 Err(unsupported_runtime(format!(
1390 "CTC alignment execution for `{}` is planned behind the alignment provider; default tests use mock alignment providers",
1391 request.model_id
1392 )))
1393 }
1394 }
1395}
1396
1397#[derive(Debug, Clone, Default)]
1399pub struct WhisperXCommandTranscriber;
1400
1401impl WhisperXCommandTranscriber {
1402 pub fn transcribe_pipeline(
1403 &mut self,
1404 request: TranscriptionPipelineRequest,
1405 ) -> Result<TranscriptionPipelineResponse> {
1406 match request.provider {
1407 TranscriptionProviderSelection::ExternalWhisperX(options) => {
1408 run_whisperx_command(request.source.path()?, options)
1409 }
1410 other => Err(invalid_request(format!(
1411 "whisperx-command cannot run provider `{}`",
1412 other.provider_id()
1413 ))),
1414 }
1415 }
1416}
1417
1418#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
1425#[serde(rename_all = "camelCase")]
1426pub struct NativeTranscriptionRunnerOptions {
1427 pub provider: TranscriptionProviderSelection,
1428 #[serde(default)]
1429 pub vad: VadOptions,
1430 #[serde(default)]
1431 pub alignment: AlignmentOptions,
1432 #[serde(default)]
1433 pub diarization: DiarizationOptions,
1434}
1435
1436impl NativeTranscriptionRunnerOptions {
1437 pub fn from_request(request: &TranscriptionPipelineRequest) -> Self {
1439 Self {
1440 provider: request.provider.clone(),
1441 vad: request.vad.clone(),
1442 alignment: request.alignment.clone(),
1443 diarization: request.diarization.clone(),
1444 }
1445 }
1446}
1447
1448impl Default for NativeTranscriptionRunnerOptions {
1449 fn default() -> Self {
1450 Self {
1451 provider: TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions::default()),
1452 vad: VadOptions::default(),
1453 alignment: AlignmentOptions::default(),
1454 diarization: DiarizationOptions::default(),
1455 }
1456 }
1457}
1458
1459pub struct NativeTranscriptionRunner {
1466 options: NativeTranscriptionRunnerOptions,
1467 mode: NativeTranscriptionRunnerMode,
1468 vad_provider: Box<dyn TranscriptionVadProvider>,
1469 asr_provider: Box<dyn AudioTranscriptionProvider>,
1470 alignment_provider: Option<Box<dyn ForcedAlignmentProvider>>,
1471 diarization_provider: Option<Box<dyn TranscriptDiarizationProvider>>,
1472}
1473
1474#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1475enum NativeTranscriptionRunnerMode {
1476 DefaultStack,
1477 CallerProviders,
1478}
1479
1480impl NativeTranscriptionRunner {
1481 pub fn new(options: NativeTranscriptionRunnerOptions) -> Result<Self> {
1488 Self::default_stack_for_options(options)
1489 }
1490
1491 fn default_stack_for_options(options: NativeTranscriptionRunnerOptions) -> Result<Self> {
1492 let asr_provider: Box<dyn AudioTranscriptionProvider> = match &options.provider {
1493 TranscriptionProviderSelection::CandleWhisper(provider_options) => Box::new(
1494 ReusableCandleWhisperTranscriber::new(provider_options.clone()),
1495 ),
1496 TranscriptionProviderSelection::WhisperCpp(provider_options) => {
1497 Box::new(WhisperCppTranscriber {
1498 options: provider_options.clone(),
1499 })
1500 }
1501 TranscriptionProviderSelection::ExternalWhisperX(_) => {
1502 return Err(invalid_request(
1503 "native transcription runner does not execute external WhisperX command providers",
1504 ));
1505 }
1506 };
1507
1508 let alignment_provider = options.alignment.enabled.then(|| {
1509 Box::new(CtcForcedAligner {
1510 options: options.alignment.clone(),
1511 }) as Box<dyn ForcedAlignmentProvider>
1512 });
1513
1514 #[cfg(feature = "diarization")]
1515 let diarization_provider = options.diarization.enabled.then(|| {
1516 Box::new(NativeSpeakerDiarizationProvider) as Box<dyn TranscriptDiarizationProvider>
1517 });
1518 #[cfg(not(feature = "diarization"))]
1519 let diarization_provider = None;
1520
1521 Ok(Self {
1522 options,
1523 mode: NativeTranscriptionRunnerMode::DefaultStack,
1524 vad_provider: Box::new(EnergyVadTranscriptionProvider),
1525 asr_provider,
1526 alignment_provider,
1527 diarization_provider,
1528 })
1529 }
1530
1531 fn rebuild_default_stack(&mut self, options: NativeTranscriptionRunnerOptions) -> Result<()> {
1532 *self = Self::default_stack_for_options(options)?;
1533 Ok(())
1534 }
1535
1536 pub fn from_providers(
1541 options: NativeTranscriptionRunnerOptions,
1542 vad_provider: Box<dyn TranscriptionVadProvider>,
1543 asr_provider: Box<dyn AudioTranscriptionProvider>,
1544 alignment_provider: Option<Box<dyn ForcedAlignmentProvider>>,
1545 diarization_provider: Option<Box<dyn TranscriptDiarizationProvider>>,
1546 ) -> Self {
1547 Self {
1548 options,
1549 mode: NativeTranscriptionRunnerMode::CallerProviders,
1550 vad_provider,
1551 asr_provider,
1552 alignment_provider,
1553 diarization_provider,
1554 }
1555 }
1556
1557 pub fn run(
1564 &mut self,
1565 request: TranscriptionPipelineRequest,
1566 observer: &mut dyn TranscriptionPipelineObserver,
1567 ) -> Result<TranscriptionPipelineResponse> {
1568 self.prepare_for_request(&request)?;
1569 let alignment_provider = self
1570 .alignment_provider
1571 .as_mut()
1572 .map(|provider| provider.as_mut() as &mut dyn ForcedAlignmentProvider);
1573 let diarization_provider = self
1574 .diarization_provider
1575 .as_mut()
1576 .map(|provider| provider.as_mut() as &mut dyn TranscriptDiarizationProvider);
1577 run_transcription_pipeline_with_observer(
1578 request,
1579 self.vad_provider.as_mut(),
1580 self.asr_provider.as_mut(),
1581 alignment_provider,
1582 diarization_provider,
1583 observer,
1584 )
1585 }
1586
1587 fn prepare_for_request(&mut self, request: &TranscriptionPipelineRequest) -> Result<()> {
1588 match self.mode {
1589 NativeTranscriptionRunnerMode::CallerProviders => {
1590 self.validate_compatible_request(request)
1591 }
1592 NativeTranscriptionRunnerMode::DefaultStack => {
1593 let request_options = NativeTranscriptionRunnerOptions::from_request(request);
1594 if request_options != self.options {
1595 self.rebuild_default_stack(request_options)?;
1596 }
1597 Ok(())
1598 }
1599 }
1600 }
1601
1602 fn validate_compatible_request(&self, request: &TranscriptionPipelineRequest) -> Result<()> {
1603 if request.provider != self.options.provider {
1604 return Err(invalid_request(
1605 "native transcription runner request provider does not match runner provider",
1606 ));
1607 }
1608 if request.vad != self.options.vad {
1609 return Err(invalid_request(
1610 "native transcription runner request VAD options do not match runner options",
1611 ));
1612 }
1613 if request.alignment != self.options.alignment {
1614 return Err(invalid_request(
1615 "native transcription runner request alignment options do not match runner options",
1616 ));
1617 }
1618 if request.diarization != self.options.diarization {
1619 return Err(invalid_request(
1620 "native transcription runner request diarization options do not match runner options",
1621 ));
1622 }
1623 Ok(())
1624 }
1625}
1626
1627fn run_native_transcription_pipeline(
1628 request: TranscriptionPipelineRequest,
1629 vad: &mut dyn TranscriptionVadProvider,
1630 asr: &mut dyn AudioTranscriptionProvider,
1631 diarization_provider: Option<&mut dyn TranscriptDiarizationProvider>,
1632) -> Result<TranscriptionPipelineResponse> {
1633 if !request.alignment.enabled {
1634 return run_transcription_pipeline(request, vad, asr, None, diarization_provider);
1635 }
1636
1637 let mut aligner = CtcForcedAligner {
1638 options: request.alignment.clone(),
1639 };
1640 run_transcription_pipeline(
1641 request,
1642 vad,
1643 asr,
1644 Some(&mut aligner as &mut dyn ForcedAlignmentProvider),
1645 diarization_provider,
1646 )
1647}
1648
1649pub fn transcribe(request: TranscriptionPipelineRequest) -> Result<TranscriptionPipelineResponse> {
1651 match &request.provider {
1652 TranscriptionProviderSelection::ExternalWhisperX(_) => {
1653 let mut provider = WhisperXCommandTranscriber;
1654 provider.transcribe_pipeline(request)
1655 }
1656 TranscriptionProviderSelection::CandleWhisper(options) => {
1657 let mut vad = EnergyVadTranscriptionProvider;
1658 let mut asr = CandleWhisperTranscriber::new(options.clone());
1659 #[cfg(feature = "diarization")]
1660 {
1661 let mut diarizer = NativeSpeakerDiarizationProvider;
1662 let diarization_provider = request
1663 .diarization
1664 .enabled
1665 .then_some(&mut diarizer as &mut dyn TranscriptDiarizationProvider);
1666 run_native_transcription_pipeline(request, &mut vad, &mut asr, diarization_provider)
1667 }
1668 #[cfg(not(feature = "diarization"))]
1669 {
1670 run_native_transcription_pipeline(request, &mut vad, &mut asr, None)
1671 }
1672 }
1673 TranscriptionProviderSelection::WhisperCpp(options) => {
1674 let mut vad = EnergyVadTranscriptionProvider;
1675 let mut asr = WhisperCppTranscriber {
1676 options: options.clone(),
1677 };
1678 #[cfg(feature = "diarization")]
1679 {
1680 let mut diarizer = NativeSpeakerDiarizationProvider;
1681 let diarization_provider = request
1682 .diarization
1683 .enabled
1684 .then_some(&mut diarizer as &mut dyn TranscriptDiarizationProvider);
1685 run_native_transcription_pipeline(request, &mut vad, &mut asr, diarization_provider)
1686 }
1687 #[cfg(not(feature = "diarization"))]
1688 {
1689 run_native_transcription_pipeline(request, &mut vad, &mut asr, None)
1690 }
1691 }
1692 }
1693}
1694
1695pub fn run_transcription_pipeline(
1697 request: TranscriptionPipelineRequest,
1698 vad_provider: &mut dyn TranscriptionVadProvider,
1699 asr_provider: &mut dyn AudioTranscriptionProvider,
1700 alignment_provider: Option<&mut dyn ForcedAlignmentProvider>,
1701 diarization_provider: Option<&mut dyn TranscriptDiarizationProvider>,
1702) -> Result<TranscriptionPipelineResponse> {
1703 let mut observer = NoopTranscriptionPipelineObserver;
1704 run_transcription_pipeline_with_observer(
1705 request,
1706 vad_provider,
1707 asr_provider,
1708 alignment_provider,
1709 diarization_provider,
1710 &mut observer,
1711 )
1712}
1713
1714pub fn run_transcription_pipeline_with_observer(
1716 request: TranscriptionPipelineRequest,
1717 vad_provider: &mut dyn TranscriptionVadProvider,
1718 asr_provider: &mut dyn AudioTranscriptionProvider,
1719 alignment_provider: Option<&mut dyn ForcedAlignmentProvider>,
1720 diarization_provider: Option<&mut dyn TranscriptDiarizationProvider>,
1721 observer: &mut dyn TranscriptionPipelineObserver,
1722) -> Result<TranscriptionPipelineResponse> {
1723 observer.observe(TranscriptionPipelineEvent::ValidationStart);
1724 ensure_pipeline_active(observer)?;
1725 validate_batch_options_for_provider(&request.provider)?;
1726 validate_task_options_for_request(&request)?;
1727 let provider = request.provider.provider_id().to_string();
1728 let model_id = request.provider.model_id().to_string();
1729 let task = request.provider.task();
1730 observer.observe(TranscriptionPipelineEvent::DecodeStart);
1731 ensure_pipeline_active(observer)?;
1732 let decode_started = Instant::now();
1733 let audio = LoadedAudio::mono_16khz_from_source(&request.source)?;
1734 observer.observe(TranscriptionPipelineEvent::DecodeEnd {
1735 duration_seconds: decode_started.elapsed().as_secs_f64(),
1736 samples: audio.samples.len(),
1737 });
1738 ensure_pipeline_active(observer)?;
1739 observer.observe(TranscriptionPipelineEvent::VadStart {
1740 provider: vad_provider.provider_id().to_string(),
1741 });
1742 ensure_pipeline_active(observer)?;
1743 let vad_response = if request.vad.enabled {
1744 vad_provider.detect_speech(VadRequest {
1745 audio: audio.clone(),
1746 options: request.vad.clone(),
1747 })?
1748 } else {
1749 VadResponse {
1750 segments: vec![SpeechActivitySegment::new(
1751 0.0,
1752 audio.duration_seconds().max(1.0 / audio.sample_rate as f64),
1753 1.0,
1754 )?],
1755 diagnostics: vec!["VAD disabled; using full source as one ASR chunk".to_string()],
1756 }
1757 };
1758 observer.observe(TranscriptionPipelineEvent::VadEnd {
1759 segments: vad_response.segments.len(),
1760 windows: diagnostic_usize(&vad_response.diagnostics, "pyannoteVadWindows")
1761 .or_else(|| diagnostic_usize(&vad_response.diagnostics, "sileroVadWindows")),
1762 });
1763 ensure_pipeline_active(observer)?;
1764
1765 let language = provider_language(&request.provider);
1766 observer.observe(TranscriptionPipelineEvent::AsrStart {
1767 model_id: model_id.clone(),
1768 });
1769 ensure_pipeline_active(observer)?;
1770 let mut asr_response = asr_provider.transcribe_with_observer(
1771 AsrRequest {
1772 audio: audio.clone(),
1773 chunks: vad_response.segments.clone(),
1774 task,
1775 language: language.clone(),
1776 model_id: model_id.clone(),
1777 },
1778 observer,
1779 )?;
1780 observer.observe(TranscriptionPipelineEvent::AsrEnd {
1781 segments: asr_response.transcript.segments.len(),
1782 });
1783 ensure_pipeline_active(observer)?;
1784 if !asr_response
1785 .diagnostics
1786 .iter()
1787 .any(|diagnostic| diagnostic.starts_with("batchChunks="))
1788 {
1789 if let TranscriptionProviderSelection::CandleWhisper(options) = &request.provider {
1790 extend_missing_candle_batch_diagnostics(
1791 &mut asr_response.diagnostics,
1792 options,
1793 vad_response.segments.len(),
1794 );
1795 }
1796 }
1797 let mut transcript = normalize_transcription_contract(asr_response.transcript)
1798 .map_err(|error| model_output_mismatch(error.to_string()))?;
1799 offset_chunk_local_segments(&mut transcript, &vad_response.segments)?;
1800
1801 let mut diagnostics = vec![format!("asrTask={}", task.as_whisper_task())];
1802 if task == TranscriptionTask::Translate {
1803 diagnostics.push("translationRuntime=whisper-task".to_string());
1804 if let Some(language) = task.output_language_hint() {
1805 diagnostics.push(format!("translationTargetLanguage={language}"));
1806 }
1807 }
1808 diagnostics.extend(vad_response.diagnostics);
1809 diagnostics.extend(asr_response.diagnostics);
1810 let mut alignment_summary = None;
1811 if request.alignment.enabled {
1812 let provider = alignment_provider.ok_or_else(|| {
1813 setup_error("alignment requested but no alignment provider is available")
1814 })?;
1815 observer.observe(TranscriptionPipelineEvent::AlignmentStart {
1816 model_id: request.alignment.model_id.clone(),
1817 });
1818 ensure_pipeline_active(observer)?;
1819 let alignment_response = provider.align_with_observer(
1820 AlignmentRequest {
1821 audio: audio.clone(),
1822 transcript: transcript.clone(),
1823 language: language.clone(),
1824 model_id: request.alignment.model_id.clone(),
1825 },
1826 observer,
1827 )?;
1828 observer.observe(TranscriptionPipelineEvent::AlignmentEnd {
1829 words: alignment_response.words.len(),
1830 });
1831 ensure_pipeline_active(observer)?;
1832 apply_alignment_words(&mut transcript, &alignment_response.words)?;
1833 apply_alignment_chars(&mut transcript, &alignment_response.chars)?;
1834 alignment_summary = Some(AlignmentSummary {
1835 provider: provider.provider_id().to_string(),
1836 model_id: alignment_response.model_id,
1837 word_count: alignment_response.words.len(),
1838 });
1839 diagnostics.extend(alignment_response.diagnostics);
1840 }
1841
1842 let mut diarization = None;
1843 if request.diarization.enabled {
1844 validate_diarization_options(&request.diarization)?;
1845 let provider = diarization_provider.ok_or_else(|| {
1846 setup_error("diarization requested but no diarization provider is available")
1847 })?;
1848 observer.observe(TranscriptionPipelineEvent::DiarizationStart {
1849 provider: diarization_progress_provider(provider.provider_id(), &request.diarization),
1850 });
1851 ensure_pipeline_active(observer)?;
1852 let response = provider.diarize(audio, &transcript, &request.diarization)?;
1853 observer.observe(TranscriptionPipelineEvent::DiarizationEnd {
1854 speakers: diarization_speaker_count(&response),
1855 segments: response.segments.len(),
1856 });
1857 ensure_pipeline_active(observer)?;
1858 diagnostics.extend(diarization_diagnostics(
1859 provider.provider_id(),
1860 &response,
1861 &request.diarization,
1862 ));
1863 diagnostics.extend(response.diagnostics.clone());
1864 transcript = audio_analysis_speakers::assign_speakers_to_transcript_with_policy(
1865 &transcript,
1866 &response,
1867 request.diarization.assignment_policy,
1868 )?;
1869 diarization = Some(response);
1870 }
1871
1872 transcript = normalize_transcription_contract(transcript)
1873 .map_err(|error| model_output_mismatch(error.to_string()))?;
1874 transcript
1875 .validate_strict()
1876 .map_err(|error| model_output_mismatch(error.to_string()))?;
1877
1878 Ok(TranscriptionPipelineResponse {
1879 accepted: true,
1880 operation: "audio.transcription.transcribe".to_string(),
1881 provider,
1882 model_id,
1883 transcript,
1884 vad_segments: vad_response.segments,
1885 alignment: alignment_summary,
1886 diarization,
1887 artifacts: Vec::new(),
1888 diagnostics,
1889 })
1890}
1891
1892fn ensure_pipeline_active(observer: &dyn TranscriptionPipelineObserver) -> Result<()> {
1893 if observer.cancellation_requested() {
1894 return Err(video_analysis_core::DetectError::InvalidArgument(
1895 "transcription cancelled at a safe workflow boundary".to_string(),
1896 ));
1897 }
1898 Ok(())
1899}
1900
1901pub fn import_whisperx_json(bytes: &[u8]) -> Result<TranscriptionContract> {
1903 text_transcripts::parse_whisperx_json(bytes)
1904 .map_err(|error| DetectError::InvalidArgument(error.to_string()))
1905}
1906
1907pub fn transcription_provider_plans() -> Vec<TranscriptionProviderPlan> {
1909 vec![
1910 candle_whisper_provider_plan(),
1911 whisper_cpp_provider_plan(),
1912 whisperx_provider_plan(),
1913 ]
1914}
1915
1916pub fn candle_whisper_provider_plan() -> TranscriptionProviderPlan {
1918 TranscriptionProviderPlan {
1919 provider_id: "candle-whisper".to_string(),
1920 external_runtime: false,
1921 wasm_supported: false,
1922 primary: true,
1923 setup: vec![
1924 "Provide an offline model bundle with config.json, generation_config.json, tokenizer.json, preprocessor_config.json, and model.safetensors.".to_string(),
1925 "Build with feature `candle`; add `cuda` for CUDA device execution.".to_string(),
1926 ],
1927 diagnostics: vec![
1928 "Candle Whisper is the primary planned Rust-native ASR and translate-to-English provider.".to_string(),
1929 "Set task=translate for native Whisper translation; wav2vec2/CTC alignment is not supported for translated output.".to_string(),
1930 "Default tests do not download models or require CUDA.".to_string(),
1931 "Default decodeRuntime=autoregressiveKvCache preserves the safe per-window KV-cache path.".to_string(),
1932 "decodeRuntime=activeRowTensorBatch enables true tensor-batched active-row decode for eligible multi-window native Candle Whisper input.".to_string(),
1933 ],
1934 }
1935}
1936
1937pub fn whisper_cpp_provider_plan() -> TranscriptionProviderPlan {
1939 TranscriptionProviderPlan {
1940 provider_id: "whisper-cpp".to_string(),
1941 external_runtime: false,
1942 wasm_supported: false,
1943 primary: false,
1944 setup: vec!["Provide a local whisper.cpp model path explicitly.".to_string()],
1945 diagnostics: vec![
1946 "whisper.cpp is retained as a native compatibility provider, not the primary ASR path."
1947 .to_string(),
1948 "Whisper translate is not supported through this provider in this crate.".to_string(),
1949 ],
1950 }
1951}
1952
1953pub fn whisperx_provider_plan() -> TranscriptionProviderPlan {
1955 TranscriptionProviderPlan {
1956 provider_id: "whisperx-command".to_string(),
1957 external_runtime: true,
1958 wasm_supported: false,
1959 primary: false,
1960 setup: vec![
1961 "Install whisperx in the active Python environment.".to_string(),
1962 "Ensure ffmpeg is available on PATH.".to_string(),
1963 "Set HF_TOKEN before diarization requests.".to_string(),
1964 ],
1965 diagnostics: vec![
1966 "WhisperX execution is opt-in and never required by default tests.".to_string(),
1967 "The compatibility command path forwards task=transcribe or task=translate to Python WhisperX.".to_string(),
1968 "Transcript normalization and WhisperX JSON import are delegated to text-transcripts."
1969 .to_string(),
1970 ],
1971 }
1972}
1973
1974fn provider_language(provider: &TranscriptionProviderSelection) -> Option<String> {
1975 match provider {
1976 TranscriptionProviderSelection::CandleWhisper(options) => options.language.clone(),
1977 TranscriptionProviderSelection::WhisperCpp(options) => options.language.clone(),
1978 TranscriptionProviderSelection::ExternalWhisperX(options) => options.language.clone(),
1979 }
1980}
1981
1982fn validate_batch_options_for_provider(provider: &TranscriptionProviderSelection) -> Result<()> {
1983 if let TranscriptionProviderSelection::CandleWhisper(options) = provider {
1984 validate_candle_batch_options(options)?;
1985 }
1986 Ok(())
1987}
1988
1989fn validate_task_options_for_request(request: &TranscriptionPipelineRequest) -> Result<()> {
1990 let task = request.provider.task();
1991 if task == TranscriptionTask::Translate && request.alignment.enabled {
1992 return Err(invalid_request(
1993 "native Whisper translation output cannot be wav2vec2/CTC-aligned against source-language audio in this implementation",
1994 ));
1995 }
1996 if matches!(
1997 request.provider,
1998 TranscriptionProviderSelection::WhisperCpp(_)
1999 ) && task == TranscriptionTask::Translate
2000 {
2001 return Err(invalid_request(
2002 "Whisper translate is not supported by the whisper.cpp provider in this crate; use candleWhisper or externalWhisperX",
2003 ));
2004 }
2005 Ok(())
2006}
2007
2008pub(crate) fn validate_candle_batch_options(options: &CandleWhisperOptions) -> Result<()> {
2009 if options.max_batch_size == Some(0) {
2010 return Err(invalid_request(
2011 "Candle Whisper max_batch_size must be greater than zero",
2012 ));
2013 }
2014 if matches!(
2015 options.decode_runtime,
2016 CandleWhisperDecodeRuntime::ActiveRowTensorBatch
2017 ) {
2018 if !options.batch_chunks {
2019 return Err(invalid_request(
2020 "Candle Whisper activeRowTensorBatch decodeRuntime requires batch_chunks=true",
2021 ));
2022 }
2023 if options.max_batch_size == Some(1) {
2024 return Err(invalid_request(
2025 "Candle Whisper activeRowTensorBatch decodeRuntime requires max_batch_size greater than one or unbounded batching",
2026 ));
2027 }
2028 }
2029 Ok(())
2030}
2031
2032pub(crate) fn candle_batch_count(options: &CandleWhisperOptions, chunk_count: usize) -> usize {
2033 if chunk_count == 0 {
2034 return 0;
2035 }
2036 if !options.batch_chunks {
2037 return chunk_count;
2038 }
2039 match options.max_batch_size {
2040 Some(max_batch_size) => chunk_count.div_ceil(max_batch_size),
2041 None => 1,
2042 }
2043}
2044
2045pub(crate) fn candle_batch_diagnostics(
2046 options: &CandleWhisperOptions,
2047 chunk_count: usize,
2048) -> Vec<String> {
2049 vec![
2050 format!("chunkCount={chunk_count}"),
2051 format!("batchChunks={}", options.batch_chunks),
2052 format!(
2053 "maxBatchSize={}",
2054 options
2055 .max_batch_size
2056 .map(|value| value.to_string())
2057 .unwrap_or_else(|| "unbounded".to_string())
2058 ),
2059 format!("batchCount={}", candle_batch_count(options, chunk_count)),
2060 format!("batchExecution={CANDLE_WHISPER_AUTOREGRESSIVE_KV_CACHE_EXECUTION}"),
2061 ]
2062}
2063
2064fn extend_missing_candle_batch_diagnostics(
2065 diagnostics: &mut Vec<String>,
2066 options: &CandleWhisperOptions,
2067 chunk_count: usize,
2068) {
2069 for diagnostic in candle_batch_diagnostics(options, chunk_count) {
2070 let Some((key, _)) = diagnostic.split_once('=') else {
2071 diagnostics.push(diagnostic);
2072 continue;
2073 };
2074 let prefix = format!("{key}=");
2075 if diagnostics.iter().any(|item| item.starts_with(&prefix)) {
2076 continue;
2077 }
2078 diagnostics.push(diagnostic);
2079 }
2080}
2081
2082pub(crate) fn validate_asr_request(request: &AsrRequest) -> Result<()> {
2083 native_audio::validate_loaded_audio(&request.audio)?;
2084 if request.chunks.is_empty() {
2085 return Err(invalid_request(
2086 "ASR request must contain at least one speech chunk",
2087 ));
2088 }
2089 let duration = request.audio.duration_seconds();
2090 let tolerance = 1.0 / request.audio.sample_rate as f64;
2091 for chunk in &request.chunks {
2092 chunk.validate()?;
2093 if chunk.end_seconds > duration + tolerance {
2094 return Err(invalid_request(format!(
2095 "speech chunk end {:.6} exceeds audio duration {:.6}",
2096 chunk.end_seconds, duration
2097 )));
2098 }
2099 }
2100 Ok(())
2101}
2102
2103fn validate_diarization_options(options: &DiarizationOptions) -> Result<()> {
2104 options.speaker.validate()
2105}
2106
2107#[cfg(feature = "diarization")]
2108fn speech_spans_from_transcript(
2109 transcript: &TranscriptionContract,
2110 audio_duration_seconds: f64,
2111) -> Result<Vec<audio_analysis_speakers::SpeechSpan>> {
2112 const AUDIO_DURATION_EPSILON: f64 = 1e-6;
2113
2114 let has_timed_words = transcript.segments.iter().any(|segment| {
2115 segment.words.iter().any(|word| {
2116 !word.text.trim().is_empty()
2117 && word.start_seconds.is_some()
2118 && word.end_seconds.is_some()
2119 })
2120 });
2121
2122 let mut spans = Vec::new();
2123 if has_timed_words {
2124 for word in transcript
2125 .segments
2126 .iter()
2127 .flat_map(|segment| &segment.words)
2128 {
2129 if word.text.trim().is_empty() {
2130 continue;
2131 }
2132 let Some((start, end)) = word.start_seconds.zip(word.end_seconds) else {
2133 continue;
2134 };
2135 spans.push(transcript_timing_span(
2136 start,
2137 end,
2138 audio_duration_seconds,
2139 AUDIO_DURATION_EPSILON,
2140 )?);
2141 }
2142 } else {
2143 for segment in &transcript.segments {
2144 if segment.text.trim().is_empty() {
2145 continue;
2146 }
2147 let Some((start, end)) = segment.start_seconds.zip(segment.end_seconds) else {
2148 continue;
2149 };
2150 spans.push(transcript_timing_span(
2151 start,
2152 end,
2153 audio_duration_seconds,
2154 AUDIO_DURATION_EPSILON,
2155 )?);
2156 }
2157 }
2158
2159 merge_transcript_speech_spans(
2160 spans,
2161 audio_analysis_speakers::EnergyVadConfig::default().merge_gap_seconds,
2162 )
2163}
2164
2165#[cfg(feature = "diarization")]
2166fn transcript_timing_span(
2167 start_seconds: f64,
2168 end_seconds: f64,
2169 audio_duration_seconds: f64,
2170 audio_duration_epsilon: f64,
2171) -> Result<audio_analysis_speakers::SpeechSpan> {
2172 if !start_seconds.is_finite() || !end_seconds.is_finite() || !audio_duration_seconds.is_finite()
2173 {
2174 return Err(invalid_request(
2175 "transcript diarization timing values must be finite",
2176 ));
2177 }
2178 if start_seconds < 0.0 || end_seconds <= start_seconds {
2179 return Err(invalid_request(
2180 "transcript diarization timing must be non-negative with positive duration",
2181 ));
2182 }
2183 if start_seconds > audio_duration_seconds + audio_duration_epsilon
2184 || end_seconds > audio_duration_seconds + audio_duration_epsilon
2185 {
2186 return Err(invalid_request(format!(
2187 "transcript diarization timing end {:.6} exceeds audio duration {:.6}",
2188 end_seconds, audio_duration_seconds
2189 )));
2190 }
2191 let end_seconds = if end_seconds > audio_duration_seconds {
2192 audio_duration_seconds
2193 } else {
2194 end_seconds
2195 };
2196 audio_analysis_speakers::SpeechSpan::new(start_seconds, end_seconds, 1.0)
2197}
2198
2199#[cfg(feature = "diarization")]
2200fn merge_transcript_speech_spans(
2201 mut spans: Vec<audio_analysis_speakers::SpeechSpan>,
2202 merge_gap_seconds: f64,
2203) -> Result<Vec<audio_analysis_speakers::SpeechSpan>> {
2204 spans.sort_by(|left, right| left.start_seconds.total_cmp(&right.start_seconds));
2205 let mut merged: Vec<audio_analysis_speakers::SpeechSpan> = Vec::new();
2206 for span in spans {
2207 if let Some(last) = merged.last_mut() {
2208 if span.start_seconds - last.end_seconds <= merge_gap_seconds {
2209 let last_duration = last.duration_seconds();
2210 let span_duration = span.duration_seconds();
2211 let total_duration = last_duration + span_duration;
2212 last.end_seconds = last.end_seconds.max(span.end_seconds);
2213 last.score = if total_duration > f64::EPSILON {
2214 (((last.score as f64 * last_duration) + (span.score as f64 * span_duration))
2215 / total_duration) as f32
2216 } else {
2217 last.score.max(span.score)
2218 };
2219 continue;
2220 }
2221 }
2222 merged.push(span);
2223 }
2224 Ok(merged)
2225}
2226
2227#[cfg(feature = "diarization")]
2228#[derive(Debug, Clone)]
2229struct TranscriptSpeechSpanVad {
2230 spans: Vec<audio_analysis_speakers::SpeechSpan>,
2231}
2232
2233#[cfg(feature = "diarization")]
2234impl audio_analysis_speakers::VoiceActivityDetector for TranscriptSpeechSpanVad {
2235 fn detect_speech(
2236 &mut self,
2237 _audio: &audio_analysis_speakers::SpeakerAudio<'_>,
2238 ) -> Result<Vec<audio_analysis_speakers::SpeechSpan>> {
2239 Ok(self.spans.clone())
2240 }
2241}
2242
2243#[cfg(feature = "diarization")]
2244fn stable_speaker_predictions_from_diarization(
2245 segments: Vec<audio_analysis_speakers::DiarizationSegment>,
2246) -> Result<Vec<SpeakerSegmentPrediction>> {
2247 let mut unknown_labels: Vec<(String, String)> = Vec::new();
2248 let mut predictions = Vec::new();
2249 for segment in segments {
2250 let speaker = match segment.speaker {
2251 audio_analysis_speakers::DiarizedSpeaker::Known(id) => id.as_str().to_string(),
2252 audio_analysis_speakers::DiarizedSpeaker::Unknown(label) => {
2253 if let Some((_, stable)) = unknown_labels
2254 .iter()
2255 .find(|(existing, _)| existing == &label)
2256 {
2257 stable.clone()
2258 } else {
2259 let stable = format!("speaker_{}", unknown_labels.len());
2260 unknown_labels.push((label, stable.clone()));
2261 stable
2262 }
2263 }
2264 };
2265 predictions.push(normalize_speaker_prediction(SpeakerSegmentPrediction {
2266 speaker,
2267 start_seconds: segment.start_seconds as f32,
2268 end_seconds: segment.end_seconds as f32,
2269 score: Some(segment.score),
2270 })?);
2271 }
2272 merge_speaker_predictions(
2273 predictions,
2274 audio_analysis_speakers::EnergyVadConfig::default().merge_gap_seconds as f32,
2275 )
2276}
2277
2278#[cfg(feature = "diarization")]
2279fn normalize_speaker_prediction(
2280 mut segment: SpeakerSegmentPrediction,
2281) -> Result<SpeakerSegmentPrediction> {
2282 segment.speaker = segment.speaker.trim().to_string();
2283 if segment.speaker.is_empty() {
2284 return Err(invalid_request("speaker label must not be empty"));
2285 }
2286 if !segment.start_seconds.is_finite() || !segment.end_seconds.is_finite() {
2287 return Err(invalid_request("speaker segment timestamps must be finite"));
2288 }
2289 if segment.end_seconds < segment.start_seconds {
2290 return Err(invalid_request(
2291 "speaker segment end_seconds must be greater than or equal to start_seconds",
2292 ));
2293 }
2294 segment.score = segment
2295 .score
2296 .and_then(|score| score.is_finite().then(|| score.clamp(0.0, 1.0)));
2297 Ok(segment)
2298}
2299
2300#[cfg(feature = "diarization")]
2301fn merge_speaker_predictions(
2302 segments: Vec<SpeakerSegmentPrediction>,
2303 merge_gap_seconds: f32,
2304) -> Result<Vec<SpeakerSegmentPrediction>> {
2305 let mut merged: Vec<SpeakerSegmentPrediction> = Vec::new();
2306 for segment in segments {
2307 if let Some(last) = merged.last_mut() {
2308 if last.speaker == segment.speaker
2309 && segment.start_seconds - last.end_seconds <= merge_gap_seconds
2310 {
2311 let last_duration = (last.end_seconds - last.start_seconds).max(0.0);
2312 let segment_duration = (segment.end_seconds - segment.start_seconds).max(0.0);
2313 let total_duration = last_duration + segment_duration;
2314 last.end_seconds = segment.end_seconds;
2315 last.score = match (last.score, segment.score) {
2316 (Some(left), Some(right)) if total_duration > f32::EPSILON => {
2317 Some(((left * last_duration) + (right * segment_duration)) / total_duration)
2318 }
2319 (Some(left), Some(right)) => Some(left.max(right)),
2320 (Some(left), None) => Some(left),
2321 (None, Some(right)) => Some(right),
2322 (None, None) => None,
2323 };
2324 continue;
2325 }
2326 }
2327 merged.push(segment);
2328 }
2329 Ok(merged)
2330}
2331
2332fn diarization_speaker_count(response: &SpeakerDiarizationResponse) -> usize {
2333 response
2334 .segments
2335 .iter()
2336 .map(|segment| segment.speaker.as_str())
2337 .collect::<std::collections::BTreeSet<_>>()
2338 .len()
2339}
2340
2341fn diagnostic_usize(diagnostics: &[String], key: &str) -> Option<usize> {
2342 let prefix = format!("{key}=");
2343 diagnostics
2344 .iter()
2345 .find_map(|diagnostic| diagnostic.strip_prefix(&prefix))
2346 .and_then(|value| value.parse().ok())
2347}
2348
2349fn diarization_progress_provider(provider_id: &str, options: &DiarizationOptions) -> String {
2350 if options
2351 .model_id
2352 .trim()
2353 .to_ascii_lowercase()
2354 .starts_with("pyannote/")
2355 {
2356 "pyannote".to_string()
2357 } else {
2358 provider_id.to_string()
2359 }
2360}
2361
2362fn diarization_diagnostics(
2363 provider_id: &str,
2364 response: &SpeakerDiarizationResponse,
2365 options: &DiarizationOptions,
2366) -> Vec<String> {
2367 let speaker_count = diarization_speaker_count(response);
2368 let mut diagnostics = vec![
2369 format!("diarizationProvider={provider_id}"),
2370 format!("diarizationRuntime={}", diarization_runtime_value(response)),
2371 format!("diarizationModelId={}", response.model_id),
2372 format!("diarizationSegmentCount={}", response.segments.len()),
2373 format!("diarizationSpeakerCount={speaker_count}"),
2374 format!(
2375 "diarizationAssignmentPolicy={}",
2376 speaker_assignment_policy_value(options.assignment_policy)
2377 ),
2378 ];
2379 if let Some(min) = options.min_speakers {
2380 diagnostics.push(format!("diarizationMinSpeakers={min}"));
2381 if speaker_count < min {
2382 diagnostics.push(format!(
2383 "diarizationSpeakerCountBelowRequestedMin={speaker_count}/{min}"
2384 ));
2385 if response.segments.len() < min {
2386 diagnostics.push("diarizationSpeakerBoundsSaturated=true".to_string());
2387 }
2388 }
2389 }
2390 if let Some(max) = options.max_speakers {
2391 diagnostics.push(format!("diarizationMaxSpeakers={max}"));
2392 if speaker_count > max {
2393 diagnostics.push(format!(
2394 "diarizationSpeakerCountAboveRequestedMax={speaker_count}/{max}"
2395 ));
2396 }
2397 }
2398 if options.min_speakers.is_some() || options.max_speakers.is_some() {
2399 diagnostics.push("diarizationSpeakerBoundsApplied=true".to_string());
2400 }
2401 if diarization_runtime_is_heuristic(response) {
2402 diagnostics.push("diarizationBaseline=heuristic-native".to_string());
2403 } else if diarization_runtime_value(response) == "onnx" {
2404 diagnostics.push("speakerEmbeddingProvider=onnx".to_string());
2405 if let Some(dimension) = options.speaker_embedding_dimension {
2406 diagnostics.push(format!("speakerEmbeddingDimension={dimension}"));
2407 }
2408 diagnostics.push("diarizationBaseline=false".to_string());
2409 }
2410 diagnostics
2411}
2412
2413fn diarization_runtime_value(response: &SpeakerDiarizationResponse) -> &'static str {
2414 match response.runtime {
2415 audio_analysis_speakers::AudioRuntime::Onnx => "onnx",
2416 audio_analysis_speakers::AudioRuntime::Candle => "candle",
2417 audio_analysis_speakers::AudioRuntime::WhisperCpp => "whisper_cpp",
2418 audio_analysis_speakers::AudioRuntime::Demucs => "demucs",
2419 audio_analysis_speakers::AudioRuntime::External => "external",
2420 audio_analysis_speakers::AudioRuntime::Spectral => "spectral",
2421 audio_analysis_speakers::AudioRuntime::Heuristic => "heuristic",
2422 audio_analysis_speakers::AudioRuntime::Imported => "imported",
2423 }
2424}
2425
2426fn diarization_runtime_is_heuristic(response: &SpeakerDiarizationResponse) -> bool {
2427 response.runtime == audio_analysis_speakers::AudioRuntime::Heuristic
2428}
2429
2430fn speaker_assignment_policy_value(policy: SpeakerAssignmentPolicy) -> &'static str {
2431 match policy {
2432 SpeakerAssignmentPolicy::Majority => "majority",
2433 SpeakerAssignmentPolicy::NearestStart => "nearestStart",
2434 SpeakerAssignmentPolicy::StrictContained => "strictContained",
2435 }
2436}
2437
2438pub(crate) fn normalize_samples_source(
2439 samples: &[f32],
2440 sample_rate: u32,
2441 channels: u16,
2442 source: Option<String>,
2443) -> Result<LoadedAudio> {
2444 if sample_rate == 0 || channels == 0 {
2445 return Err(DetectError::InvalidAudioFormat {
2446 sample_rate,
2447 channels,
2448 });
2449 }
2450 if samples.is_empty() {
2451 return Err(invalid_request("empty audio"));
2452 }
2453 if !samples.len().is_multiple_of(channels as usize) {
2454 return Err(invalid_request(
2455 "sample count must contain complete interleaved frames",
2456 ));
2457 }
2458 if samples.iter().any(|sample| !sample.is_finite()) {
2459 return Err(invalid_request("audio samples must be finite"));
2460 }
2461 let mono = if channels == 1 {
2462 samples.to_vec()
2463 } else {
2464 samples
2465 .chunks_exact(channels as usize)
2466 .map(|frame| frame.iter().sum::<f32>() / channels as f32)
2467 .collect::<Vec<_>>()
2468 };
2469 let samples = if sample_rate == 16_000 {
2470 mono
2471 } else {
2472 resample_linear(&mono, sample_rate, 16_000)
2473 };
2474 Ok(LoadedAudio {
2475 samples,
2476 sample_rate: 16_000,
2477 channels: 1,
2478 source,
2479 })
2480}
2481
2482pub(crate) fn resample_linear(samples: &[f32], from_rate: u32, to_rate: u32) -> Vec<f32> {
2483 if samples.is_empty() || from_rate == to_rate {
2484 return samples.to_vec();
2485 }
2486 let output_len = ((samples.len() as f64 * to_rate as f64) / from_rate as f64)
2487 .round()
2488 .max(1.0) as usize;
2489 (0..output_len)
2490 .map(|index| {
2491 let position = index as f64 * from_rate as f64 / to_rate as f64;
2492 let left = position.floor() as usize;
2493 let right = (left + 1).min(samples.len() - 1);
2494 let frac = (position - left as f64) as f32;
2495 samples[left] * (1.0 - frac) + samples[right] * frac
2496 })
2497 .collect()
2498}
2499
2500fn energy_vad_segments(
2501 audio: &LoadedAudio,
2502 options: &VadOptions,
2503) -> Result<Vec<SpeechActivitySegment>> {
2504 validate_vad_options(options)?;
2505 let frame_size = seconds_to_samples(options.frame_seconds, audio.sample_rate)?;
2506 let hop_size = seconds_to_samples(options.hop_seconds, audio.sample_rate)?;
2507 let mut active = Vec::new();
2508 let mut start = 0;
2509 while start < audio.samples.len() {
2510 let end = (start + frame_size).min(audio.samples.len());
2511 let score = rms(&audio.samples[start..end]);
2512 if score >= options.rms_threshold {
2513 active.push(SpeechActivitySegment::new(
2514 start as f64 / audio.sample_rate as f64,
2515 end as f64 / audio.sample_rate as f64,
2516 score,
2517 )?);
2518 }
2519 if start + hop_size >= audio.samples.len() {
2520 break;
2521 }
2522 start += hop_size;
2523 }
2524 let duration = audio.duration_seconds();
2525 let mut merged = merge_speech_segments(active, options.merge_gap_seconds)?;
2526 for segment in &mut merged {
2527 segment.start_seconds = (segment.start_seconds - options.padding_seconds).max(0.0);
2528 segment.end_seconds = (segment.end_seconds + options.padding_seconds).min(duration);
2529 }
2530 merged = merge_speech_segments(merged, options.merge_gap_seconds)?;
2531 let filtered = merged
2532 .into_iter()
2533 .filter(|segment| segment.end_seconds - segment.start_seconds >= options.min_speech_seconds)
2534 .flat_map(|segment| split_max_chunk(segment, options.max_chunk_seconds))
2535 .collect::<Vec<_>>();
2536 if filtered.is_empty() {
2537 return Ok(vec![SpeechActivitySegment::new(
2538 0.0,
2539 duration.max(1.0 / audio.sample_rate as f64),
2540 0.0,
2541 )?]);
2542 }
2543 Ok(filtered)
2544}
2545
2546fn validate_vad_options(options: &VadOptions) -> Result<()> {
2547 if !options.rms_threshold.is_finite() || options.rms_threshold < 0.0 {
2548 return Err(invalid_request(
2549 "VAD RMS threshold must be finite and non-negative",
2550 ));
2551 }
2552 for (name, value, positive) in [
2553 ("frameSeconds", options.frame_seconds, true),
2554 ("hopSeconds", options.hop_seconds, true),
2555 ("minSpeechSeconds", options.min_speech_seconds, false),
2556 ("paddingSeconds", options.padding_seconds, false),
2557 ("mergeGapSeconds", options.merge_gap_seconds, false),
2558 ("maxChunkSeconds", options.max_chunk_seconds, true),
2559 ] {
2560 if !value.is_finite() || (positive && value <= 0.0) || (!positive && value < 0.0) {
2561 return Err(invalid_request(format!(
2562 "VAD option `{name}` has an invalid value"
2563 )));
2564 }
2565 }
2566 Ok(())
2567}
2568
2569fn seconds_to_samples(seconds: f64, sample_rate: u32) -> Result<usize> {
2570 if !seconds.is_finite() || seconds <= 0.0 || sample_rate == 0 {
2571 return Err(invalid_request("invalid sample duration"));
2572 }
2573 Ok((seconds * sample_rate as f64).round().max(1.0) as usize)
2574}
2575
2576fn rms(samples: &[f32]) -> f32 {
2577 if samples.is_empty() {
2578 return 0.0;
2579 }
2580 (samples.iter().map(|sample| sample * sample).sum::<f32>() / samples.len() as f32).sqrt()
2581}
2582
2583fn merge_speech_segments(
2584 mut segments: Vec<SpeechActivitySegment>,
2585 merge_gap_seconds: f64,
2586) -> Result<Vec<SpeechActivitySegment>> {
2587 segments.sort_by(|left, right| left.start_seconds.total_cmp(&right.start_seconds));
2588 let mut merged: Vec<SpeechActivitySegment> = Vec::new();
2589 for segment in segments {
2590 if let Some(last) = merged.last_mut() {
2591 if segment.start_seconds - last.end_seconds <= merge_gap_seconds {
2592 last.end_seconds = last.end_seconds.max(segment.end_seconds);
2593 last.score = last.score.max(segment.score);
2594 continue;
2595 }
2596 }
2597 merged.push(segment);
2598 }
2599 for segment in &merged {
2600 segment.validate()?;
2601 }
2602 Ok(merged)
2603}
2604
2605fn split_max_chunk(
2606 segment: SpeechActivitySegment,
2607 max_chunk_seconds: f64,
2608) -> Vec<SpeechActivitySegment> {
2609 if segment.end_seconds - segment.start_seconds <= max_chunk_seconds {
2610 return vec![segment];
2611 }
2612 let mut chunks = Vec::new();
2613 let mut start = segment.start_seconds;
2614 while start < segment.end_seconds {
2615 let end = (start + max_chunk_seconds).min(segment.end_seconds);
2616 if end > start {
2617 chunks.push(SpeechActivitySegment {
2618 start_seconds: start,
2619 end_seconds: end,
2620 score: segment.score,
2621 });
2622 }
2623 start = end;
2624 }
2625 chunks
2626}
2627
2628fn offset_chunk_local_segments(
2629 transcript: &mut TranscriptionContract,
2630 chunks: &[SpeechActivitySegment],
2631) -> Result<()> {
2632 if chunks.is_empty() || transcript.segments.len() != chunks.len() {
2633 return Ok(());
2634 }
2635 for (segment, chunk) in transcript.segments.iter_mut().zip(chunks) {
2636 if segment
2637 .attributes
2638 .get("timing")
2639 .is_some_and(|value| value == "global")
2640 {
2641 continue;
2642 }
2643 if let Some(start) = &mut segment.start_seconds {
2644 *start += chunk.start_seconds;
2645 }
2646 if let Some(end) = &mut segment.end_seconds {
2647 *end += chunk.start_seconds;
2648 }
2649 for word in &mut segment.words {
2650 if let Some(start) = &mut word.start_seconds {
2651 *start += chunk.start_seconds;
2652 }
2653 if let Some(end) = &mut word.end_seconds {
2654 *end += chunk.start_seconds;
2655 }
2656 }
2657 segment
2658 .attributes
2659 .insert("timing".to_string(), "global".to_string());
2660 }
2661 Ok(())
2662}
2663
2664fn apply_alignment_words(
2665 transcript: &mut TranscriptionContract,
2666 words: &[AlignedWord],
2667) -> Result<()> {
2668 for aligned in words {
2669 if !aligned.start_seconds.is_finite()
2670 || !aligned.end_seconds.is_finite()
2671 || aligned.end_seconds < aligned.start_seconds
2672 {
2673 return Err(model_output_mismatch(
2674 "alignment output contains invalid word timing",
2675 ));
2676 }
2677 let segment = transcript
2678 .segments
2679 .iter_mut()
2680 .find(|segment| segment.index == aligned.segment_index)
2681 .ok_or_else(|| model_output_mismatch("alignment output references unknown segment"))?;
2682 while segment.words.len() <= aligned.word_index {
2683 segment.words.push(TranscriptWordContract {
2684 text: String::new(),
2685 start_seconds: None,
2686 end_seconds: None,
2687 confidence: None,
2688 speaker: None,
2689 attributes: BTreeMap::new(),
2690 });
2691 }
2692 let word = &mut segment.words[aligned.word_index];
2693 if word.text.trim().is_empty() {
2694 word.text = aligned.text.clone();
2695 }
2696 word.start_seconds = Some(aligned.start_seconds);
2697 word.end_seconds = Some(aligned.end_seconds);
2698 word.confidence = aligned.confidence;
2699 }
2700 for segment in &mut transcript.segments {
2701 let timed_words = segment
2702 .words
2703 .iter()
2704 .filter_map(|word| word.start_seconds.zip(word.end_seconds))
2705 .collect::<Vec<_>>();
2706 if let (Some((start, _)), Some((_, end))) = (timed_words.first(), timed_words.last()) {
2707 segment.start_seconds = Some(*start);
2708 segment.end_seconds = Some(*end);
2709 }
2710 }
2711 Ok(())
2712}
2713
2714fn apply_alignment_chars(
2715 transcript: &mut TranscriptionContract,
2716 chars: &[AlignedChar],
2717) -> Result<()> {
2718 for aligned in chars {
2719 if let Some((start, end)) = aligned.start_seconds.zip(aligned.end_seconds) {
2720 if !start.is_finite() || !end.is_finite() || end < start {
2721 return Err(model_output_mismatch(
2722 "alignment output contains invalid char timing",
2723 ));
2724 }
2725 }
2726 let segment = transcript
2727 .segments
2728 .iter_mut()
2729 .find(|segment| segment.index == aligned.segment_index)
2730 .ok_or_else(|| model_output_mismatch("alignment output references unknown segment"))?;
2731 while segment.chars.len() <= aligned.char_index {
2732 segment.chars.push(TranscriptCharContract {
2733 character: String::new(),
2734 start_seconds: None,
2735 end_seconds: None,
2736 confidence: None,
2737 attributes: BTreeMap::new(),
2738 });
2739 }
2740 let character = &mut segment.chars[aligned.char_index];
2741 if character.character.is_empty() {
2742 character.character = aligned.character.clone();
2743 }
2744 character.start_seconds = aligned.start_seconds;
2745 character.end_seconds = aligned.end_seconds;
2746 character.confidence = aligned.confidence;
2747 }
2748 Ok(())
2749}
2750
2751fn validate_candle_setup(options: &CandleWhisperOptions) -> Result<()> {
2752 validate_candle_batch_options(options)?;
2753 let resolved_device = native_device::resolve_native_device(options.device)?;
2754 options
2755 .compute_type
2756 .resolve_for_device(resolved_device.cuda_active())?;
2757 if !cfg!(feature = "candle") {
2758 return Err(unsupported_runtime(format!(
2759 "Candle Whisper requested but the binary lacks the `candle` feature; {}; build with `candle` for native execution and `model-bundles` for Hugging Face cache resolution",
2760 candle_whisper_setup_context(options)
2761 )));
2762 }
2763 if let Some(bundle) = &options.model_bundle {
2764 validate_model_bundle_files(
2765 bundle,
2766 &[
2767 "config.json",
2768 "generation_config.json",
2769 "tokenizer.json",
2770 "preprocessor_config.json",
2771 "model.safetensors",
2772 ],
2773 )
2774 } else {
2775 Ok(())
2776 }
2777}
2778
2779#[cfg(not(feature = "alignment"))]
2780fn validate_alignment_setup(options: &AlignmentOptions) -> Result<()> {
2781 if !cfg!(feature = "alignment") {
2782 return Err(unsupported_runtime(
2783 "CTC alignment requested but the binary lacks the `alignment` feature",
2784 ));
2785 }
2786 if let Some(bundle) = &options.model_bundle {
2787 validate_model_bundle_files(
2788 bundle,
2789 &[
2790 "config.json",
2791 "tokenizer.json",
2792 "preprocessor_config.json",
2793 "model.safetensors",
2794 ],
2795 )
2796 } else {
2797 Err(setup_error(
2798 "required CTC alignment model bundle is missing",
2799 ))
2800 }
2801}
2802
2803fn validate_model_bundle_files(bundle: &Path, files: &[&str]) -> Result<()> {
2804 native_bundles::validate_required_bundle_files(bundle, files)
2805}
2806
2807fn run_whisperx_command(
2808 source_path: &Path,
2809 options: WhisperXCommandOptions,
2810) -> Result<TranscriptionPipelineResponse> {
2811 let task = options.task;
2812 let output_dir = options
2813 .output_dir
2814 .clone()
2815 .unwrap_or_else(default_whisperx_output_dir);
2816 fs::create_dir_all(&output_dir)?;
2817
2818 let hf_token = if options.diarize {
2819 let env_name = options
2820 .hf_token_env
2821 .clone()
2822 .unwrap_or_else(|| "HF_TOKEN".to_string());
2823 Some(std::env::var(&env_name).map_err(|_| {
2824 setup_error(format!(
2825 "diarization requires `{env_name}` to be set before running WhisperX"
2826 ))
2827 })?)
2828 } else {
2829 None
2830 };
2831
2832 let args = whisperx_args(source_path, &output_dir, &options, hf_token.as_deref());
2833 let child = Command::new(&options.command)
2834 .args(&args)
2835 .stdin(Stdio::null())
2836 .stdout(Stdio::piped())
2837 .stderr(Stdio::piped())
2838 .spawn()
2839 .map_err(|error| {
2840 if error.kind() == std::io::ErrorKind::NotFound {
2841 setup_error(format!(
2842 "WhisperX command `{}` was not found; install whisperx or pass provider.command",
2843 options.command.display()
2844 ))
2845 } else {
2846 DetectError::Io(error)
2847 }
2848 })?;
2849 let output = wait_with_optional_timeout(child, &options.command, options.timeout_seconds)?;
2850 if !output.status.success() {
2851 return Err(setup_error(format!(
2852 "WhisperX command `{}` failed: {}",
2853 options.command.display(),
2854 String::from_utf8_lossy(&output.stderr).trim()
2855 )));
2856 }
2857
2858 let (transcript_path, transcript_bytes) =
2859 whisperx_json_bytes(source_path, &output_dir, &output.stdout).ok_or_else(|| {
2860 model_output_mismatch(format!(
2861 "WhisperX completed but no JSON transcript was found in `{}`",
2862 output_dir.display()
2863 ))
2864 })?;
2865 let mut transcript = import_whisperx_json(&transcript_bytes)?;
2866 if transcript.source.is_none() {
2867 transcript.source = Some(source_path.to_string_lossy().into_owned());
2868 }
2869 let vad_segments = whisperx_stdout_vad_segments(&output.stdout);
2870 let artifacts = transcript_path
2871 .map(|path| TranscriptionArtifact {
2872 kind: "whisperx-json".to_string(),
2873 path,
2874 })
2875 .into_iter()
2876 .collect();
2877 let mut diagnostics = vec![
2878 format!("asrTask={}", task.as_whisper_task()),
2879 format!("ran WhisperX output in `{}`", output_dir.display()),
2880 "parsed WhisperX JSON through text-transcripts".to_string(),
2881 ];
2882 if task == TranscriptionTask::Translate {
2883 diagnostics.push("translationRuntime=whisperx-command".to_string());
2884 if let Some(language) = task.output_language_hint() {
2885 diagnostics.push(format!("translationTargetLanguage={language}"));
2886 }
2887 }
2888 if !vad_segments.is_empty() {
2889 diagnostics.push(format!(
2890 "whisperxVadSegmentsFromStdout={}",
2891 vad_segments.len()
2892 ));
2893 }
2894 Ok(TranscriptionPipelineResponse {
2895 accepted: true,
2896 operation: "audio.transcription.transcribe".to_string(),
2897 provider: "whisperx-command".to_string(),
2898 model_id: options.model,
2899 transcript,
2900 vad_segments,
2901 alignment: None,
2902 diarization: None,
2903 artifacts,
2904 diagnostics,
2905 })
2906}
2907
2908fn whisperx_args(
2909 source_path: &Path,
2910 output_dir: &Path,
2911 options: &WhisperXCommandOptions,
2912 hf_token: Option<&str>,
2913) -> Vec<String> {
2914 let mut args = vec![
2915 source_path.to_string_lossy().into_owned(),
2916 "--model".to_string(),
2917 options.model.clone(),
2918 "--task".to_string(),
2919 options.task.as_whisper_task().to_string(),
2920 "--device".to_string(),
2921 options.device.as_str().to_string(),
2922 "--output_format".to_string(),
2923 "json".to_string(),
2924 "--output_dir".to_string(),
2925 output_dir.to_string_lossy().into_owned(),
2926 ];
2927 if let Some(language) = &options.language {
2928 args.extend(["--language".to_string(), language.clone()]);
2929 }
2930 if let Some(compute_type) = &options.compute_type {
2931 args.extend(["--compute_type".to_string(), compute_type.clone()]);
2932 }
2933 if let Some(batch_size) = options.batch_size {
2934 args.extend(["--batch_size".to_string(), batch_size.to_string()]);
2935 }
2936 if options.no_align {
2937 args.push("--no_align".to_string());
2938 }
2939 if let Some(align_model) = &options.align_model {
2940 args.extend(["--align_model".to_string(), align_model.clone()]);
2941 }
2942 if let Some(model_dir) = &options.model_dir {
2943 args.extend([
2944 "--model_dir".to_string(),
2945 model_dir.to_string_lossy().into_owned(),
2946 ]);
2947 }
2948 if options.model_cache_only {
2949 args.push("--model_cache_only".to_string());
2950 }
2951 args.extend([
2952 "--interpolate_method".to_string(),
2953 options.interpolate_method.as_whisperx_arg().to_string(),
2954 ]);
2955 if options.return_char_alignments {
2956 args.push("--return_char_alignments".to_string());
2957 }
2958 if options.diarize {
2959 args.push("--diarize".to_string());
2960 if let Some(min_speakers) = options.min_speakers {
2961 args.extend(["--min_speakers".to_string(), min_speakers.to_string()]);
2962 }
2963 if let Some(max_speakers) = options.max_speakers {
2964 args.extend(["--max_speakers".to_string(), max_speakers.to_string()]);
2965 }
2966 if let Some(hf_token) = hf_token {
2967 args.extend(["--hf_token".to_string(), hf_token.to_string()]);
2968 }
2969 }
2970 args.extend(options.extra_args.clone());
2971 args
2972}
2973
2974fn whisperx_json_bytes(
2975 source_path: &Path,
2976 output_dir: &Path,
2977 stdout: &[u8],
2978) -> Option<(Option<PathBuf>, Vec<u8>)> {
2979 let path = find_json_artifact_for_source(output_dir, source_path);
2980 if let Some(path) = path {
2981 return fs::read(&path).ok().map(|bytes| (Some(path), bytes));
2982 }
2983 serde_json::from_slice::<serde_json::Value>(stdout)
2984 .ok()
2985 .map(|_| (None, stdout.to_vec()))
2986}
2987
2988fn whisperx_stdout_vad_segments(stdout: &[u8]) -> Vec<SpeechActivitySegment> {
2989 let stdout = String::from_utf8_lossy(stdout);
2990 stdout
2991 .lines()
2992 .filter_map(|line| {
2993 let start_marker = "Transcript: [";
2994 let start = line.find(start_marker)? + start_marker.len();
2995 let range = line[start..].split_once(']')?.0;
2996 let (start_seconds, end_seconds) = range.split_once("-->")?;
2997 let start_seconds = start_seconds.trim().parse::<f64>().ok()?;
2998 let end_seconds = end_seconds.trim().parse::<f64>().ok()?;
2999 SpeechActivitySegment::new(start_seconds, end_seconds, 1.0).ok()
3000 })
3001 .collect()
3002}
3003
3004fn find_json_artifact_for_source(output_dir: &Path, source_path: &Path) -> Option<PathBuf> {
3005 let expected = source_path
3006 .file_stem()
3007 .and_then(|stem| stem.to_str())
3008 .filter(|stem| !stem.is_empty())
3009 .map(|stem| output_dir.join(format!("{stem}.json")));
3010 if let Some(expected) = expected.filter(|path| path.is_file()) {
3011 return Some(expected);
3012 }
3013
3014 let mut candidates = fs::read_dir(output_dir)
3015 .ok()?
3016 .filter_map(|entry| entry.ok().map(|entry| entry.path()))
3017 .filter(|path| path.extension().and_then(|value| value.to_str()) == Some("json"))
3018 .collect::<Vec<_>>();
3019 candidates.sort();
3020 candidates.into_iter().next()
3021}
3022
3023fn default_whisperx_output_dir() -> PathBuf {
3024 let millis = SystemTime::now()
3025 .duration_since(UNIX_EPOCH)
3026 .unwrap_or_default()
3027 .as_millis();
3028 std::env::temp_dir().join(format!("video-analysis-whisperx-{millis}"))
3029}
3030
3031fn wait_with_optional_timeout(
3032 mut child: Child,
3033 command: &Path,
3034 timeout_seconds: Option<u64>,
3035) -> Result<Output> {
3036 let Some(seconds) = timeout_seconds else {
3037 return child.wait_with_output().map_err(DetectError::Io);
3038 };
3039 let started = Instant::now();
3040 loop {
3041 if child.try_wait()?.is_some() {
3042 return child.wait_with_output().map_err(DetectError::Io);
3043 }
3044 if started.elapsed() >= Duration::from_secs(seconds) {
3045 let _ = child.kill();
3046 let _ = child.wait();
3047 return Err(timeout_error(format!(
3048 "WhisperX command `{}` timed out after {seconds} seconds",
3049 command.display()
3050 )));
3051 }
3052 std::thread::sleep(Duration::from_millis(25));
3053 }
3054}
3055
3056pub(crate) fn setup_error(message: impl Into<String>) -> DetectError {
3057 DetectError::InvalidArgument(format!("setup_error: {}", message.into()))
3058}
3059
3060pub(crate) fn invalid_request(message: impl Into<String>) -> DetectError {
3061 DetectError::InvalidArgument(format!("invalid_request: {}", message.into()))
3062}
3063
3064pub(crate) fn model_output_mismatch(message: impl Into<String>) -> DetectError {
3065 DetectError::InvalidArgument(format!("model_output_mismatch: {}", message.into()))
3066}
3067
3068fn timeout_error(message: impl Into<String>) -> DetectError {
3069 DetectError::InvalidArgument(format!("timeout: {}", message.into()))
3070}
3071
3072pub(crate) fn unsupported_runtime(message: impl Into<String>) -> DetectError {
3073 DetectError::InvalidArgument(format!("unsupported_runtime: {}", message.into()))
3074}
3075
3076#[cfg(test)]
3077mod tests {
3078 use super::*;
3079 use text_transcripts::TranscriptSegmentContract;
3080
3081 #[derive(Default)]
3082 struct MockAsrProvider;
3083
3084 impl AudioTranscriptionProvider for MockAsrProvider {
3085 fn provider_id(&self) -> &str {
3086 "mock-asr"
3087 }
3088
3089 fn transcribe(&mut self, request: AsrRequest) -> Result<AsrResponse> {
3090 let segments = request
3091 .chunks
3092 .iter()
3093 .enumerate()
3094 .map(|(index, chunk)| {
3095 let mut segment = TranscriptSegmentContract::new(index as u64, " hello ");
3096 segment.start_seconds = Some(0.0);
3097 segment.end_seconds = Some(chunk.end_seconds - chunk.start_seconds);
3098 segment
3099 })
3100 .collect::<Vec<_>>();
3101 Ok(AsrResponse {
3102 model_id: request.model_id,
3103 language: request
3104 .task
3105 .output_language_hint()
3106 .map(str::to_string)
3107 .or(request.language),
3108 transcript: TranscriptionContract::from_segments(
3109 request.audio.source,
3110 Some("en".to_string()),
3111 segments,
3112 )
3113 .map_err(|error| DetectError::InvalidArgument(error.to_string()))?,
3114 diagnostics: vec!["mock ASR completed".to_string()],
3115 })
3116 }
3117 }
3118
3119 #[derive(Clone)]
3120 struct FixedVadProvider {
3121 segments: Vec<SpeechActivitySegment>,
3122 }
3123
3124 impl TranscriptionVadProvider for FixedVadProvider {
3125 fn provider_id(&self) -> &str {
3126 "fixed-vad"
3127 }
3128
3129 fn detect_speech(&mut self, _request: VadRequest) -> Result<VadResponse> {
3130 Ok(VadResponse {
3131 segments: self.segments.clone(),
3132 diagnostics: vec!["fixed VAD completed".to_string()],
3133 })
3134 }
3135 }
3136
3137 #[derive(Default)]
3138 struct MockAlignmentProvider;
3139
3140 impl ForcedAlignmentProvider for MockAlignmentProvider {
3141 fn provider_id(&self) -> &str {
3142 "mock-aligner"
3143 }
3144
3145 fn align(&mut self, request: AlignmentRequest) -> Result<AlignmentResponse> {
3146 Ok(AlignmentResponse {
3147 model_id: request.model_id,
3148 words: vec![AlignedWord {
3149 segment_index: 0,
3150 word_index: 0,
3151 text: "hello".to_string(),
3152 start_seconds: 0.05,
3153 end_seconds: 0.35,
3154 confidence: Some(0.91),
3155 }],
3156 chars: Vec::new(),
3157 diagnostics: vec!["mock alignment completed".to_string()],
3158 })
3159 }
3160 }
3161
3162 #[derive(Default)]
3163 struct RecordingObserver {
3164 events: Vec<TranscriptionPipelineEvent>,
3165 }
3166
3167 impl TranscriptionPipelineObserver for RecordingObserver {
3168 fn observe(&mut self, event: TranscriptionPipelineEvent) {
3169 self.events.push(event);
3170 }
3171 }
3172
3173 #[derive(Default)]
3174 struct CancelAtAlignmentObserver {
3175 cancelled: bool,
3176 }
3177
3178 impl TranscriptionPipelineObserver for CancelAtAlignmentObserver {
3179 fn observe(&mut self, event: TranscriptionPipelineEvent) {
3180 if matches!(event, TranscriptionPipelineEvent::AlignmentStart { .. }) {
3181 self.cancelled = true;
3182 }
3183 }
3184
3185 fn cancellation_requested(&self) -> bool {
3186 self.cancelled
3187 }
3188 }
3189
3190 struct ObservingAsrProvider;
3191
3192 impl AudioTranscriptionProvider for ObservingAsrProvider {
3193 fn provider_id(&self) -> &str {
3194 "observing-asr"
3195 }
3196
3197 fn transcribe(&mut self, request: AsrRequest) -> Result<AsrResponse> {
3198 MockAsrProvider.transcribe(request)
3199 }
3200
3201 fn transcribe_with_observer(
3202 &mut self,
3203 request: AsrRequest,
3204 observer: &mut dyn TranscriptionPipelineObserver,
3205 ) -> Result<AsrResponse> {
3206 observer.observe(TranscriptionPipelineEvent::ModelLoadStart {
3207 stage: "asr".to_string(),
3208 provider: self.provider_id().to_string(),
3209 model_id: request.model_id.clone(),
3210 });
3211 observer.observe(TranscriptionPipelineEvent::ModelLoadEnd {
3212 stage: "asr".to_string(),
3213 provider: self.provider_id().to_string(),
3214 model_id: request.model_id.clone(),
3215 duration_seconds: 0.125,
3216 });
3217 observer.observe(TranscriptionPipelineEvent::ModelReuse {
3218 stage: "asr".to_string(),
3219 provider: self.provider_id().to_string(),
3220 model_id: request.model_id.clone(),
3221 });
3222 self.transcribe(request)
3223 }
3224 }
3225
3226 struct ObservingAlignmentProvider;
3227
3228 impl ForcedAlignmentProvider for ObservingAlignmentProvider {
3229 fn provider_id(&self) -> &str {
3230 "observing-aligner"
3231 }
3232
3233 fn align(&mut self, request: AlignmentRequest) -> Result<AlignmentResponse> {
3234 MockAlignmentProvider.align(request)
3235 }
3236
3237 fn align_with_observer(
3238 &mut self,
3239 request: AlignmentRequest,
3240 observer: &mut dyn TranscriptionPipelineObserver,
3241 ) -> Result<AlignmentResponse> {
3242 observer.observe(TranscriptionPipelineEvent::ModelLoadStart {
3243 stage: "alignment".to_string(),
3244 provider: self.provider_id().to_string(),
3245 model_id: request.model_id.clone(),
3246 });
3247 observer.observe(TranscriptionPipelineEvent::ModelLoadEnd {
3248 stage: "alignment".to_string(),
3249 provider: self.provider_id().to_string(),
3250 model_id: request.model_id.clone(),
3251 duration_seconds: 0.25,
3252 });
3253 self.align(request)
3254 }
3255 }
3256
3257 struct MockDiarizationProvider;
3258
3259 impl TranscriptDiarizationProvider for MockDiarizationProvider {
3260 fn provider_id(&self) -> &str {
3261 "mock-diarization"
3262 }
3263
3264 fn diarize(
3265 &mut self,
3266 _audio: LoadedAudio,
3267 _transcript: &TranscriptionContract,
3268 options: &DiarizationOptions,
3269 ) -> Result<SpeakerDiarizationResponse> {
3270 Ok(SpeakerDiarizationResponse {
3271 accepted: true,
3272 operation: "audio.speakers.diarize".to_string(),
3273 model_id: options.model_id.clone(),
3274 runtime: audio_analysis_speakers::AudioRuntime::Imported,
3275 segments: vec![SpeakerSegmentPrediction {
3276 speaker: "SPEAKER_00".to_string(),
3277 start_seconds: 0.0,
3278 end_seconds: 1.0,
3279 score: Some(0.9),
3280 }],
3281 speaker_embeddings: None,
3282 diagnostics: Vec::new(),
3283 })
3284 }
3285 }
3286
3287 struct PanickingDiarizationProvider {
3288 called: bool,
3289 }
3290
3291 #[cfg(feature = "diarization")]
3292 struct MockOnnxDiarizationProvider;
3293
3294 #[cfg(feature = "diarization")]
3295 impl TranscriptDiarizationProvider for MockOnnxDiarizationProvider {
3296 fn provider_id(&self) -> &str {
3297 "mock-onnx-diarization"
3298 }
3299
3300 fn diarize(
3301 &mut self,
3302 _audio: LoadedAudio,
3303 _transcript: &TranscriptionContract,
3304 options: &DiarizationOptions,
3305 ) -> Result<SpeakerDiarizationResponse> {
3306 Ok(SpeakerDiarizationResponse {
3307 accepted: true,
3308 operation: "audio.speakers.diarize".to_string(),
3309 model_id: options.model_id.clone(),
3310 runtime: audio_analysis_speakers::AudioRuntime::Onnx,
3311 segments: vec![audio_analysis_speakers::SpeakerSegmentPrediction {
3312 speaker: "SPEAKER_ONNX".to_string(),
3313 start_seconds: 0.0,
3314 end_seconds: 1.0,
3315 score: Some(0.9),
3316 }],
3317 speaker_embeddings: None,
3318 diagnostics: Vec::new(),
3319 })
3320 }
3321 }
3322
3323 impl TranscriptDiarizationProvider for PanickingDiarizationProvider {
3324 fn provider_id(&self) -> &str {
3325 "panicking-diarization"
3326 }
3327
3328 fn diarize(
3329 &mut self,
3330 _audio: LoadedAudio,
3331 _transcript: &TranscriptionContract,
3332 _options: &DiarizationOptions,
3333 ) -> Result<SpeakerDiarizationResponse> {
3334 self.called = true;
3335 panic!("diarization provider should not be called for invalid options");
3336 }
3337 }
3338
3339 fn sample_request() -> TranscriptionPipelineRequest {
3340 let mut samples = vec![0.0; 16_000];
3341 for sample in &mut samples[1_000..5_000] {
3342 *sample = 0.1;
3343 }
3344 TranscriptionPipelineRequest {
3345 source: TranscriptionSource::Samples {
3346 samples,
3347 sample_rate: 16_000,
3348 channels: 1,
3349 source: Some("synthetic".to_string()),
3350 },
3351 provider: TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions::default()),
3352 vad: VadOptions {
3353 min_speech_seconds: 0.01,
3354 ..VadOptions::default()
3355 },
3356 alignment: AlignmentOptions::default(),
3357 diarization: DiarizationOptions::default(),
3358 output: TranscriptionOutputOptions::default(),
3359 }
3360 }
3361
3362 #[test]
3363 fn cooperative_cancellation_stops_before_alignment_provider_work() {
3364 let mut request = sample_request();
3365 request.alignment.enabled = true;
3366 let mut vad = FixedVadProvider {
3367 segments: vec![SpeechActivitySegment::new(0.0, 1.0, 1.0).unwrap()],
3368 };
3369 let mut asr = MockAsrProvider;
3370 let mut aligner = MockAlignmentProvider;
3371 let mut observer = CancelAtAlignmentObserver::default();
3372
3373 let error = run_transcription_pipeline_with_observer(
3374 request,
3375 &mut vad,
3376 &mut asr,
3377 Some(&mut aligner),
3378 None,
3379 &mut observer,
3380 )
3381 .expect_err("cancellation should win before alignment provider work");
3382
3383 assert!(error.to_string().contains("cancelled"));
3384 }
3385
3386 fn batch_test_chunks() -> Vec<SpeechActivitySegment> {
3387 vec![
3388 SpeechActivitySegment::new(0.0, 0.20, 1.0).unwrap(),
3389 SpeechActivitySegment::new(0.20, 0.40, 1.0).unwrap(),
3390 SpeechActivitySegment::new(0.40, 0.60, 1.0).unwrap(),
3391 ]
3392 }
3393
3394 fn batch_test_request(options: CandleWhisperOptions) -> TranscriptionPipelineRequest {
3395 TranscriptionPipelineRequest {
3396 provider: TranscriptionProviderSelection::CandleWhisper(options),
3397 ..sample_request()
3398 }
3399 }
3400
3401 #[test]
3402 fn pipeline_observer_receives_model_load_and_reuse_events() {
3403 let mut request = sample_request();
3404 request.vad = VadOptions {
3405 enabled: false,
3406 ..VadOptions::default()
3407 };
3408 request.alignment = AlignmentOptions {
3409 enabled: true,
3410 model_id: "facebook/wav2vec2-base-960h".to_string(),
3411 ..AlignmentOptions::default()
3412 };
3413 let mut vad = FixedVadProvider {
3414 segments: vec![SpeechActivitySegment::new(0.0, 0.5, 1.0).unwrap()],
3415 };
3416 let mut asr = ObservingAsrProvider;
3417 let mut aligner = ObservingAlignmentProvider;
3418 let mut observer = RecordingObserver::default();
3419
3420 let response = run_transcription_pipeline_with_observer(
3421 request,
3422 &mut vad,
3423 &mut asr,
3424 Some(&mut aligner),
3425 None,
3426 &mut observer,
3427 )
3428 .expect("pipeline should run with observing providers");
3429
3430 assert!(response.accepted);
3431 assert!(observer
3432 .events
3433 .contains(&TranscriptionPipelineEvent::ModelLoadStart {
3434 stage: "asr".to_string(),
3435 provider: "observing-asr".to_string(),
3436 model_id: "openai/whisper-large-v3-turbo".to_string(),
3437 }));
3438 assert!(observer
3439 .events
3440 .contains(&TranscriptionPipelineEvent::ModelLoadEnd {
3441 stage: "asr".to_string(),
3442 provider: "observing-asr".to_string(),
3443 model_id: "openai/whisper-large-v3-turbo".to_string(),
3444 duration_seconds: 0.125,
3445 }));
3446 assert!(observer
3447 .events
3448 .contains(&TranscriptionPipelineEvent::ModelReuse {
3449 stage: "asr".to_string(),
3450 provider: "observing-asr".to_string(),
3451 model_id: "openai/whisper-large-v3-turbo".to_string(),
3452 }));
3453 assert!(observer
3454 .events
3455 .contains(&TranscriptionPipelineEvent::ModelLoadStart {
3456 stage: "alignment".to_string(),
3457 provider: "observing-aligner".to_string(),
3458 model_id: "facebook/wav2vec2-base-960h".to_string(),
3459 }));
3460 assert!(observer
3461 .events
3462 .contains(&TranscriptionPipelineEvent::ModelLoadEnd {
3463 stage: "alignment".to_string(),
3464 provider: "observing-aligner".to_string(),
3465 model_id: "facebook/wav2vec2-base-960h".to_string(),
3466 duration_seconds: 0.25,
3467 }));
3468 }
3469
3470 #[test]
3471 fn candle_whisper_options_default_to_automatic_compute_type() {
3472 assert_eq!(
3473 CandleWhisperOptions::default().compute_type,
3474 CandleWhisperComputeType::Automatic
3475 );
3476 }
3477
3478 #[test]
3479 fn candle_whisper_compute_type_serializes_public_values_and_aliases() {
3480 let options = CandleWhisperOptions {
3481 compute_type: CandleWhisperComputeType::Fp16,
3482 ..CandleWhisperOptions::default()
3483 };
3484 let encoded = serde_json::to_value(&options).unwrap();
3485 assert_eq!(encoded["computeType"], "fp16");
3486
3487 let decoded: CandleWhisperOptions =
3488 serde_json::from_value(serde_json::json!({"computeType": "float32"})).unwrap();
3489 assert_eq!(decoded.compute_type, CandleWhisperComputeType::Fp32);
3490
3491 let decoded: CandleWhisperOptions =
3492 serde_json::from_value(serde_json::json!({"computeType": "auto"})).unwrap();
3493 assert_eq!(decoded.compute_type, CandleWhisperComputeType::Automatic);
3494 }
3495
3496 #[test]
3497 fn candle_whisper_compute_type_resolves_by_device() {
3498 assert_eq!(
3499 CandleWhisperComputeType::Automatic
3500 .resolve_for_device(true)
3501 .unwrap(),
3502 CandleWhisperComputeType::Fp16
3503 );
3504 assert_eq!(
3505 CandleWhisperComputeType::Fp16
3506 .resolve_for_device(true)
3507 .unwrap(),
3508 CandleWhisperComputeType::Fp16
3509 );
3510 assert_eq!(
3511 CandleWhisperComputeType::Fp32
3512 .resolve_for_device(true)
3513 .unwrap(),
3514 CandleWhisperComputeType::Fp32
3515 );
3516 assert_eq!(
3517 CandleWhisperComputeType::Automatic
3518 .resolve_for_device(false)
3519 .unwrap(),
3520 CandleWhisperComputeType::Fp32
3521 );
3522 assert_eq!(
3523 CandleWhisperComputeType::Fp32
3524 .resolve_for_device(false)
3525 .unwrap(),
3526 CandleWhisperComputeType::Fp32
3527 );
3528 }
3529
3530 #[test]
3531 fn candle_whisper_cpu_fp16_is_rejected_clearly() {
3532 let error = CandleWhisperComputeType::Fp16
3533 .resolve_for_device(false)
3534 .unwrap_err()
3535 .to_string();
3536
3537 assert!(error.contains("setup_error"));
3538 assert!(error.contains("fp16 requires a CUDA device"));
3539 }
3540
3541 #[test]
3542 fn only_automatic_candle_compute_type_is_setup_fallback_eligible() {
3543 assert!(CandleWhisperComputeType::Automatic.setup_fallback_eligible());
3544 assert!(!CandleWhisperComputeType::Fp16.setup_fallback_eligible());
3545 assert!(!CandleWhisperComputeType::Fp32.setup_fallback_eligible());
3546 }
3547
3548 fn diarization_response_for_tests(
3549 segments: Vec<SpeakerSegmentPrediction>,
3550 ) -> SpeakerDiarizationResponse {
3551 SpeakerDiarizationResponse {
3552 accepted: true,
3553 operation: "audio.speakers.diarize".to_string(),
3554 model_id: "test-speakers".to_string(),
3555 runtime: audio_analysis_speakers::AudioRuntime::Imported,
3556 segments,
3557 speaker_embeddings: None,
3558 diagnostics: Vec::new(),
3559 }
3560 }
3561
3562 fn transcript_with_words(
3563 words: Vec<(&str, f64, f64)>,
3564 ) -> std::result::Result<TranscriptionContract, Box<dyn std::error::Error>> {
3565 let mut segment = TranscriptSegmentContract::new(
3566 0,
3567 words
3568 .iter()
3569 .map(|(word, _, _)| *word)
3570 .collect::<Vec<_>>()
3571 .join(" "),
3572 );
3573 segment.start_seconds = Some(0.0);
3574 segment.end_seconds = Some(2.0);
3575 segment.words = words
3576 .into_iter()
3577 .map(
3578 |(text, start_seconds, end_seconds)| TranscriptWordContract {
3579 text: text.to_string(),
3580 start_seconds: Some(start_seconds),
3581 end_seconds: Some(end_seconds),
3582 confidence: None,
3583 speaker: None,
3584 attributes: BTreeMap::new(),
3585 },
3586 )
3587 .collect();
3588 Ok(TranscriptionContract::from_segments(
3589 None,
3590 Some("en".to_string()),
3591 vec![segment],
3592 )?)
3593 }
3594
3595 fn assign_speakers_from_diarization(
3596 transcript: &mut TranscriptionContract,
3597 diarization: &SpeakerDiarizationResponse,
3598 policy: SpeakerAssignmentPolicy,
3599 ) -> Result<()> {
3600 *transcript = audio_analysis_speakers::assign_speakers_to_transcript_with_policy(
3601 transcript,
3602 diarization,
3603 policy,
3604 )?;
3605 Ok(())
3606 }
3607
3608 #[cfg(feature = "diarization")]
3609 fn sine_into(samples: &mut [f32], sample_rate: u32, start_seconds: f32, freq_hz: f32) {
3610 for (offset, sample) in samples.iter_mut().enumerate() {
3611 let t = start_seconds + offset as f32 / sample_rate as f32;
3612 *sample = (2.0 * std::f32::consts::PI * freq_hz * t).sin() * 0.5;
3613 }
3614 }
3615
3616 #[cfg(feature = "diarization")]
3617 fn two_profile_loaded_audio() -> LoadedAudio {
3618 let sample_rate = 16_000;
3619 let mut samples = vec![0.0_f32; sample_rate as usize * 2];
3620 let first_start = (0.20 * sample_rate as f32) as usize;
3621 let first_end = (0.50 * sample_rate as f32) as usize;
3622 let second_start = (1.00 * sample_rate as f32) as usize;
3623 let second_end = (1.40 * sample_rate as f32) as usize;
3624 sine_into(
3625 &mut samples[first_start..first_end],
3626 sample_rate,
3627 0.20,
3628 220.0,
3629 );
3630 sine_into(
3631 &mut samples[second_start..second_end],
3632 sample_rate,
3633 1.00,
3634 1_200.0,
3635 );
3636 LoadedAudio {
3637 samples,
3638 sample_rate,
3639 channels: 1,
3640 source: Some("synthetic-two-speaker".to_string()),
3641 }
3642 }
3643
3644 #[cfg(all(feature = "diarization", feature = "onnx"))]
3645 fn local_16khz_wav_samples(
3646 path: &Path,
3647 ) -> std::result::Result<Vec<f32>, Box<dyn std::error::Error>> {
3648 let mut reader = hound::WavReader::open(path)?;
3649 let spec = reader.spec();
3650 if spec.sample_rate != 16_000 {
3651 return Err(format!(
3652 "DIARIZATION_AUDIO_PATH must be 16 kHz WAV, got {} Hz",
3653 spec.sample_rate
3654 )
3655 .into());
3656 }
3657 if spec.channels == 0 {
3658 return Err("DIARIZATION_AUDIO_PATH WAV channel count must be non-zero".into());
3659 }
3660 let interleaved = match spec.sample_format {
3661 hound::SampleFormat::Float => {
3662 if spec.bits_per_sample != 32 {
3663 return Err("float DIARIZATION_AUDIO_PATH WAV must be 32-bit".into());
3664 }
3665 reader
3666 .samples::<f32>()
3667 .collect::<std::result::Result<Vec<_>, _>>()?
3668 }
3669 hound::SampleFormat::Int if spec.bits_per_sample <= 16 => reader
3670 .samples::<i16>()
3671 .map(|sample| sample.map(|value| value as f32 / 32_768.0))
3672 .collect::<std::result::Result<Vec<_>, _>>()?,
3673 hound::SampleFormat::Int => {
3674 let scale = 2_f32.powi(spec.bits_per_sample as i32 - 1);
3675 reader
3676 .samples::<i32>()
3677 .map(|sample| sample.map(|value| value as f32 / scale))
3678 .collect::<std::result::Result<Vec<_>, _>>()?
3679 }
3680 };
3681 let channels = spec.channels as usize;
3682 let samples = if channels == 1 {
3683 interleaved
3684 } else {
3685 interleaved
3686 .chunks_exact(channels)
3687 .map(|frame| frame.iter().copied().sum::<f32>() / channels as f32)
3688 .collect()
3689 };
3690 if samples.is_empty() {
3691 return Err("DIARIZATION_AUDIO_PATH WAV must contain samples".into());
3692 }
3693 if !samples.iter().all(|sample| sample.is_finite()) {
3694 return Err("DIARIZATION_AUDIO_PATH WAV samples must be finite".into());
3695 }
3696 Ok(samples)
3697 }
3698
3699 #[test]
3700 fn diarization_options_reject_invalid_speaker_bounds() {
3701 let mut options = DiarizationOptions {
3702 speaker: SpeakerDiarizationOptions {
3703 min_speakers: Some(0),
3704 ..SpeakerDiarizationOptions::default()
3705 },
3706 ..DiarizationOptions::default()
3707 };
3708 assert!(validate_diarization_options(&options)
3709 .unwrap_err()
3710 .to_string()
3711 .contains("invalid_request"));
3712
3713 options = DiarizationOptions {
3714 speaker: SpeakerDiarizationOptions {
3715 max_speakers: Some(0),
3716 ..SpeakerDiarizationOptions::default()
3717 },
3718 ..DiarizationOptions::default()
3719 };
3720 assert!(validate_diarization_options(&options)
3721 .unwrap_err()
3722 .to_string()
3723 .contains("invalid_request"));
3724
3725 options = DiarizationOptions {
3726 speaker: SpeakerDiarizationOptions {
3727 min_speakers: Some(3),
3728 max_speakers: Some(2),
3729 ..SpeakerDiarizationOptions::default()
3730 },
3731 ..DiarizationOptions::default()
3732 };
3733 assert!(validate_diarization_options(&options)
3734 .unwrap_err()
3735 .to_string()
3736 .contains("invalid_request"));
3737
3738 options = DiarizationOptions {
3739 speaker: SpeakerDiarizationOptions {
3740 min_speakers: Some(1),
3741 max_speakers: Some(2),
3742 ..SpeakerDiarizationOptions::default()
3743 },
3744 ..DiarizationOptions::default()
3745 };
3746 validate_diarization_options(&options).unwrap();
3747 }
3748
3749 #[test]
3750 fn diarization_options_keep_flat_json_shape_with_speaker_owned_contract() {
3751 let options: DiarizationOptions = serde_json::from_value(serde_json::json!({
3752 "enabled": true,
3753 "modelId": "native-speakers",
3754 "speakerEmbeddingModelBundle": "/models/speaker",
3755 "speakerEmbeddingModelFile": "speaker.onnx",
3756 "speakerEmbeddingInputName": "waveform",
3757 "speakerEmbeddingOutputName": "embedding",
3758 "speakerEmbeddingDimension": 192,
3759 "speakerEmbeddingSampleRate": 16000,
3760 "returnSpeakerEmbeddings": true,
3761 "minSpeakers": 1,
3762 "maxSpeakers": 2,
3763 "assignmentPolicy": "strictContained"
3764 }))
3765 .unwrap();
3766
3767 assert!(options.enabled);
3768 assert_eq!(options.model_id, "native-speakers");
3769 assert_eq!(
3770 options.speaker_embedding_model_bundle.as_deref(),
3771 Some(Path::new("/models/speaker"))
3772 );
3773 assert_eq!(
3774 options.speaker_embedding_model_file.as_deref(),
3775 Some("speaker.onnx")
3776 );
3777 assert_eq!(
3778 options.speaker_embedding_input_name.as_deref(),
3779 Some("waveform")
3780 );
3781 assert_eq!(
3782 options.speaker_embedding_output_name.as_deref(),
3783 Some("embedding")
3784 );
3785 assert_eq!(options.speaker_embedding_dimension, Some(192));
3786 assert_eq!(options.speaker_embedding_sample_rate, Some(16_000));
3787 assert!(options.return_speaker_embeddings);
3788 assert_eq!(options.min_speakers, Some(1));
3789 assert_eq!(options.max_speakers, Some(2));
3790 assert_eq!(
3791 options.assignment_policy,
3792 SpeakerAssignmentPolicy::StrictContained
3793 );
3794
3795 let value = serde_json::to_value(&options).unwrap();
3796 assert_eq!(value["enabled"], true);
3797 assert_eq!(value["modelId"], "native-speakers");
3798 assert_eq!(value["assignmentPolicy"], "strictContained");
3799 assert!(value.get("speaker").is_none());
3800 }
3801
3802 #[test]
3803 fn default_build_exposes_speakers_owned_diarization_response_types() {
3804 fn accepts_speakers_response(_: audio_analysis_speakers::SpeakerDiarizationResponse) {}
3805
3806 let response = SpeakerDiarizationResponse {
3807 accepted: true,
3808 operation: "audio.speakers.diarize".to_string(),
3809 model_id: "fixture".to_string(),
3810 runtime: audio_analysis_speakers::AudioRuntime::Imported,
3811 segments: vec![SpeakerSegmentPrediction {
3812 speaker: "speaker_0".to_string(),
3813 start_seconds: 0.0,
3814 end_seconds: 1.0,
3815 score: None,
3816 }],
3817 speaker_embeddings: None,
3818 diagnostics: Vec::new(),
3819 };
3820
3821 accepts_speakers_response(response);
3822 }
3823
3824 #[test]
3825 fn batch_options_reject_zero_max_batch_size() {
3826 let mut vad = FixedVadProvider {
3827 segments: batch_test_chunks(),
3828 };
3829 let mut asr = MockAsrProvider;
3830 let result = run_transcription_pipeline(
3831 batch_test_request(CandleWhisperOptions {
3832 max_batch_size: Some(0),
3833 ..CandleWhisperOptions::default()
3834 }),
3835 &mut vad,
3836 &mut asr,
3837 None,
3838 None,
3839 );
3840
3841 let error = result.unwrap_err().to_string();
3842 assert!(error.contains("invalid_request"));
3843 assert!(error.contains("max_batch_size"));
3844 }
3845
3846 #[test]
3847 fn candle_decode_runtime_defaults_to_autoregressive_kv_cache() {
3848 let options = CandleWhisperOptions::default();
3849
3850 assert_eq!(
3851 options.decode_runtime,
3852 CandleWhisperDecodeRuntime::AutoregressiveKvCache
3853 );
3854 assert_eq!(
3855 options.decode_runtime.execution_id(),
3856 "candle-whisper-autoregressive-kv-cache"
3857 );
3858 validate_candle_batch_options(&options).unwrap();
3859 }
3860
3861 #[test]
3862 fn candle_decode_runtime_deserializes_active_row_tensor_batch() {
3863 let options: CandleWhisperOptions = serde_json::from_value(serde_json::json!({
3864 "decodeRuntime": "activeRowTensorBatch",
3865 "batchChunks": true,
3866 "maxBatchSize": 4
3867 }))
3868 .unwrap();
3869
3870 assert_eq!(
3871 options.decode_runtime,
3872 CandleWhisperDecodeRuntime::ActiveRowTensorBatch
3873 );
3874 assert_eq!(
3875 options.decode_runtime.execution_id(),
3876 "candle-whisper-active-row-tensor-batch"
3877 );
3878 }
3879
3880 #[test]
3881 fn candle_active_row_decode_runtime_is_supported_when_batching_is_enabled() {
3882 let options = CandleWhisperOptions {
3883 decode_runtime: CandleWhisperDecodeRuntime::ActiveRowTensorBatch,
3884 batch_chunks: true,
3885 max_batch_size: Some(4),
3886 ..CandleWhisperOptions::default()
3887 };
3888
3889 validate_candle_batch_options(&options).unwrap();
3890 assert!(options.decode_runtime.is_supported());
3891 }
3892
3893 #[test]
3894 fn candle_active_row_decode_runtime_requires_chunk_batching() {
3895 let options = CandleWhisperOptions {
3896 decode_runtime: CandleWhisperDecodeRuntime::ActiveRowTensorBatch,
3897 batch_chunks: false,
3898 max_batch_size: Some(4),
3899 ..CandleWhisperOptions::default()
3900 };
3901
3902 let error = validate_candle_batch_options(&options)
3903 .unwrap_err()
3904 .to_string();
3905 assert!(error.contains("invalid_request"));
3906 assert!(error.contains("requires batch_chunks=true"));
3907 }
3908
3909 #[test]
3910 fn candle_active_row_decode_runtime_rejects_single_row_batching() {
3911 let options = CandleWhisperOptions {
3912 decode_runtime: CandleWhisperDecodeRuntime::ActiveRowTensorBatch,
3913 batch_chunks: true,
3914 max_batch_size: Some(1),
3915 ..CandleWhisperOptions::default()
3916 };
3917
3918 let error = validate_candle_batch_options(&options)
3919 .unwrap_err()
3920 .to_string();
3921 assert!(error.contains("invalid_request"));
3922 assert!(error.contains("max_batch_size greater than one"));
3923 }
3924
3925 #[test]
3926 fn batch_chunking_preserves_transcript_order() {
3927 let mut vad = FixedVadProvider {
3928 segments: batch_test_chunks(),
3929 };
3930 let mut asr = MockAsrProvider;
3931 let response = run_transcription_pipeline(
3932 batch_test_request(CandleWhisperOptions {
3933 batch_chunks: true,
3934 max_batch_size: Some(2),
3935 ..CandleWhisperOptions::default()
3936 }),
3937 &mut vad,
3938 &mut asr,
3939 None,
3940 None,
3941 )
3942 .unwrap();
3943
3944 let starts = response
3945 .transcript
3946 .segments
3947 .iter()
3948 .map(|segment| segment.start_seconds.unwrap())
3949 .collect::<Vec<_>>();
3950 assert_eq!(starts, vec![0.0, 0.20, 0.40]);
3951 }
3952
3953 #[test]
3954 fn batch_chunking_reports_batch_diagnostics() {
3955 let mut vad = FixedVadProvider {
3956 segments: batch_test_chunks(),
3957 };
3958 let mut asr = MockAsrProvider;
3959 let response = run_transcription_pipeline(
3960 batch_test_request(CandleWhisperOptions {
3961 batch_chunks: true,
3962 max_batch_size: Some(2),
3963 ..CandleWhisperOptions::default()
3964 }),
3965 &mut vad,
3966 &mut asr,
3967 None,
3968 None,
3969 )
3970 .unwrap();
3971
3972 assert!(response
3973 .diagnostics
3974 .iter()
3975 .any(|item| item == "chunkCount=3"));
3976 assert!(response
3977 .diagnostics
3978 .iter()
3979 .any(|item| item == "batchChunks=true"));
3980 assert!(response
3981 .diagnostics
3982 .iter()
3983 .any(|item| item == "maxBatchSize=2"));
3984 assert!(response
3985 .diagnostics
3986 .iter()
3987 .any(|item| item == "batchCount=2"));
3988 assert!(response
3989 .diagnostics
3990 .iter()
3991 .any(|item| item == "batchExecution=candle-whisper-autoregressive-kv-cache"));
3992 }
3993
3994 #[test]
3995 fn requested_active_row_runtime_reports_fallback_execution_without_native_proof() {
3996 let mut vad = FixedVadProvider {
3997 segments: batch_test_chunks(),
3998 };
3999 let mut asr = MockAsrProvider;
4000 let response = run_transcription_pipeline(
4001 batch_test_request(CandleWhisperOptions {
4002 decode_runtime: CandleWhisperDecodeRuntime::ActiveRowTensorBatch,
4003 batch_chunks: true,
4004 max_batch_size: Some(3),
4005 ..CandleWhisperOptions::default()
4006 }),
4007 &mut vad,
4008 &mut asr,
4009 None,
4010 None,
4011 )
4012 .unwrap();
4013
4014 assert!(response
4015 .diagnostics
4016 .iter()
4017 .any(|item| item == "batchExecution=candle-whisper-autoregressive-kv-cache"));
4018 assert!(!response
4019 .diagnostics
4020 .iter()
4021 .any(|item| item == "batchExecution=candle-whisper-active-row-tensor-batch"));
4022 }
4023
4024 #[test]
4025 fn public_batch_diagnostics_do_not_claim_active_row_execution_by_request() {
4026 let diagnostics = candle_batch_diagnostics(
4027 &CandleWhisperOptions {
4028 decode_runtime: CandleWhisperDecodeRuntime::ActiveRowTensorBatch,
4029 batch_chunks: true,
4030 max_batch_size: Some(4),
4031 ..CandleWhisperOptions::default()
4032 },
4033 3,
4034 );
4035
4036 assert!(diagnostics
4037 .iter()
4038 .any(|item| item == "batchExecution=candle-whisper-autoregressive-kv-cache"));
4039 assert!(!diagnostics
4040 .iter()
4041 .any(|item| item == "batchExecution=candle-whisper-active-row-tensor-batch"));
4042 }
4043
4044 #[test]
4045 fn batch_chunking_reports_unbounded_batch_diagnostics() {
4046 let mut vad = FixedVadProvider {
4047 segments: batch_test_chunks(),
4048 };
4049 let mut asr = MockAsrProvider;
4050 let response = run_transcription_pipeline(
4051 batch_test_request(CandleWhisperOptions {
4052 batch_chunks: true,
4053 max_batch_size: None,
4054 ..CandleWhisperOptions::default()
4055 }),
4056 &mut vad,
4057 &mut asr,
4058 None,
4059 None,
4060 )
4061 .unwrap();
4062
4063 assert!(response
4064 .diagnostics
4065 .iter()
4066 .any(|item| item == "maxBatchSize=unbounded"));
4067 assert!(response
4068 .diagnostics
4069 .iter()
4070 .any(|item| item == "batchCount=1"));
4071 }
4072
4073 #[test]
4074 fn batch_disabled_reports_sequential_diagnostics() {
4075 let mut vad = FixedVadProvider {
4076 segments: batch_test_chunks(),
4077 };
4078 let mut asr = MockAsrProvider;
4079 let response = run_transcription_pipeline(
4080 batch_test_request(CandleWhisperOptions {
4081 batch_chunks: false,
4082 max_batch_size: Some(2),
4083 ..CandleWhisperOptions::default()
4084 }),
4085 &mut vad,
4086 &mut asr,
4087 None,
4088 None,
4089 )
4090 .unwrap();
4091
4092 assert!(response
4093 .diagnostics
4094 .iter()
4095 .any(|item| item == "batchChunks=false"));
4096 assert!(response
4097 .diagnostics
4098 .iter()
4099 .any(|item| item == "batchCount=3"));
4100 assert!(response
4101 .diagnostics
4102 .iter()
4103 .any(|item| item == "batchExecution=candle-whisper-autoregressive-kv-cache"));
4104 }
4105
4106 #[test]
4107 #[cfg(feature = "diarization")]
4108 fn speech_spans_from_transcript_prefers_aligned_words(
4109 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4110 let transcript = transcript_with_words(vec![("hello", 0.2, 0.5), ("world", 1.0, 1.4)])?;
4111
4112 let spans = speech_spans_from_transcript(&transcript, 2.0)?;
4113
4114 assert_eq!(spans.len(), 2);
4115 assert_eq!(spans[0].start_seconds, 0.2);
4116 assert_eq!(spans[0].end_seconds, 0.5);
4117 assert_eq!(spans[1].start_seconds, 1.0);
4118 assert_eq!(spans[1].end_seconds, 1.4);
4119 assert!(spans[0].start_seconds < spans[1].start_seconds);
4120 assert!(spans
4121 .iter()
4122 .all(|span| span.score.is_finite() && span.score > 0.0));
4123 Ok(())
4124 }
4125
4126 #[test]
4127 #[cfg(feature = "diarization")]
4128 fn speech_spans_from_transcript_falls_back_to_segments(
4129 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4130 let mut first = TranscriptSegmentContract::new(0, "hello");
4131 first.start_seconds = Some(0.2);
4132 first.end_seconds = Some(0.5);
4133 let mut second = TranscriptSegmentContract::new(1, "world");
4134 second.start_seconds = Some(1.0);
4135 second.end_seconds = Some(1.4);
4136 let transcript = TranscriptionContract::from_segments(
4137 None,
4138 Some("en".to_string()),
4139 vec![second, first],
4140 )?;
4141
4142 let spans = speech_spans_from_transcript(&transcript, 2.0)?;
4143
4144 assert_eq!(spans.len(), 2);
4145 assert_eq!(spans[0].start_seconds, 0.2);
4146 assert_eq!(spans[0].end_seconds, 0.5);
4147 assert_eq!(spans[1].start_seconds, 1.0);
4148 assert_eq!(spans[1].end_seconds, 1.4);
4149 Ok(())
4150 }
4151
4152 #[test]
4153 #[cfg(feature = "diarization")]
4154 fn speech_spans_from_transcript_rejects_out_of_range_timing(
4155 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4156 let transcript = transcript_with_words(vec![("hello", 0.2, 2.000_01)])?;
4157
4158 let error = speech_spans_from_transcript(&transcript, 2.0)
4159 .unwrap_err()
4160 .to_string();
4161
4162 assert!(error.contains("invalid_request"));
4163 Ok(())
4164 }
4165
4166 #[cfg(feature = "alignment")]
4167 struct AlignmentAwareDiarizationProvider {
4168 saw_aligned_words: bool,
4169 }
4170
4171 #[cfg(feature = "alignment")]
4172 impl TranscriptDiarizationProvider for AlignmentAwareDiarizationProvider {
4173 fn provider_id(&self) -> &str {
4174 "alignment-aware-diarization"
4175 }
4176
4177 fn diarize(
4178 &mut self,
4179 _audio: LoadedAudio,
4180 transcript: &TranscriptionContract,
4181 options: &DiarizationOptions,
4182 ) -> Result<SpeakerDiarizationResponse> {
4183 self.saw_aligned_words = transcript.segments.iter().any(|segment| {
4184 segment.words.iter().any(|word| {
4185 word.start_seconds.is_some()
4186 && word.end_seconds.is_some()
4187 && word.confidence.is_some()
4188 })
4189 });
4190 assert!(
4191 self.saw_aligned_words,
4192 "diarization should receive transcript word timings from alignment"
4193 );
4194 let mut response = diarization_response_for_tests(vec![SpeakerSegmentPrediction {
4195 speaker: "SPEAKER_ALIGNED".to_string(),
4196 start_seconds: 0.0,
4197 end_seconds: 1.0,
4198 score: Some(0.95),
4199 }]);
4200 response.model_id = options.model_id.clone();
4201 Ok(response)
4202 }
4203 }
4204
4205 #[cfg(all(feature = "alignment", feature = "candle"))]
4206 fn write_tiny_wav2vec2_bundle(root: &Path) {
4207 use candle_core::{Device, Tensor};
4208 use std::collections::HashMap;
4209
4210 fs::write(
4211 root.join("config.json"),
4212 serde_json::json!({
4213 "model_type": "wav2vec2",
4214 "architectures": ["Wav2Vec2ForCTC"],
4215 "vocab_size": 10,
4216 "word_delimiter_token": "|",
4217 "hidden_size": 1,
4218 "num_hidden_layers": 0,
4219 "num_attention_heads": 1,
4220 "intermediate_size": 1,
4221 "hidden_act": "gelu",
4222 "layer_norm_eps": 1e-5,
4223 "feat_extract_activation": "gelu",
4224 "conv_dim": [1],
4225 "conv_stride": [1],
4226 "conv_kernel": [1],
4227 "conv_bias": false,
4228 "num_conv_pos_embeddings": 0,
4229 "num_conv_pos_embedding_groups": 1
4230 })
4231 .to_string(),
4232 )
4233 .unwrap();
4234 fs::write(
4235 root.join("tokenizer.json"),
4236 serde_json::json!({
4237 "version": "1.0",
4238 "word_delimiter_token": "|",
4239 "model": {
4240 "type": "WordLevel",
4241 "vocab": {
4242 "[PAD]": 0,
4243 "H": 1,
4244 "E": 2,
4245 "L": 3,
4246 "O": 4,
4247 "|": 5,
4248 "W": 6,
4249 "R": 7,
4250 "D": 8,
4251 "<unk>": 9
4252 },
4253 "unk_token": "<unk>"
4254 }
4255 })
4256 .to_string(),
4257 )
4258 .unwrap();
4259 fs::write(
4260 root.join("preprocessor_config.json"),
4261 serde_json::json!({
4262 "sampling_rate": 16000,
4263 "do_normalize": false,
4264 "return_attention_mask": false
4265 })
4266 .to_string(),
4267 )
4268 .unwrap();
4269
4270 let device = Device::Cpu;
4271 let mut tensors = HashMap::new();
4272 tensors.insert(
4273 "wav2vec2.feature_extractor.conv_layers.0.conv.weight".to_string(),
4274 Tensor::new(&[1.0f32], &device)
4275 .unwrap()
4276 .reshape((1, 1, 1))
4277 .unwrap(),
4278 );
4279 tensors.insert(
4280 "wav2vec2.feature_projection.layer_norm.weight".to_string(),
4281 Tensor::new(&[1.0f32], &device).unwrap(),
4282 );
4283 tensors.insert(
4284 "wav2vec2.feature_projection.layer_norm.bias".to_string(),
4285 Tensor::new(&[0.0f32], &device).unwrap(),
4286 );
4287 tensors.insert(
4288 "wav2vec2.feature_projection.projection.weight".to_string(),
4289 Tensor::new(&[1.0f32], &device)
4290 .unwrap()
4291 .reshape((1, 1))
4292 .unwrap(),
4293 );
4294 tensors.insert(
4295 "wav2vec2.feature_projection.projection.bias".to_string(),
4296 Tensor::new(&[0.0f32], &device).unwrap(),
4297 );
4298 tensors.insert(
4299 "lm_head.weight".to_string(),
4300 Tensor::new(&[0.0f32; 10], &device)
4301 .unwrap()
4302 .reshape((10, 1))
4303 .unwrap(),
4304 );
4305 tensors.insert(
4306 "lm_head.bias".to_string(),
4307 Tensor::new(&[0.0f32; 10], &device).unwrap(),
4308 );
4309 candle_core::safetensors::save(&tensors, root.join("model.safetensors")).unwrap();
4310 }
4311
4312 #[cfg(all(feature = "alignment", feature = "candle", feature = "model-bundles"))]
4313 fn env_path(name: &str) -> Option<PathBuf> {
4314 std::env::var_os(name).map(PathBuf::from)
4315 }
4316
4317 #[cfg(all(feature = "alignment", feature = "candle", feature = "model-bundles"))]
4318 fn write_default_alignment_smoke_wav(path: &Path) -> std::result::Result<(), hound::Error> {
4319 let spec = hound::WavSpec {
4320 channels: 1,
4321 sample_rate: 16_000,
4322 bits_per_sample: 16,
4323 sample_format: hound::SampleFormat::Int,
4324 };
4325 let mut writer = hound::WavWriter::create(path, spec)?;
4326 for index in 0..16_000 {
4327 let sample = if (1_000..8_000).contains(&index) {
4328 16_384i16
4329 } else {
4330 0i16
4331 };
4332 writer.write_sample(sample)?;
4333 }
4334 writer.finalize()
4335 }
4336
4337 #[cfg(all(feature = "alignment", feature = "candle", feature = "model-bundles"))]
4338 fn default_alignment_smoke_audio_path(
4339 temp: &Path,
4340 ) -> std::result::Result<PathBuf, Box<dyn std::error::Error>> {
4341 if let Some(path) =
4342 env_path("ALIGNMENT_AUDIO_PATH").or_else(|| env_path("TRANSCRIPTION_AUDIO_PATH"))
4343 {
4344 return Ok(path);
4345 }
4346
4347 let repo_sample = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
4348 .join("../../../vendor/whisper.cpp/samples/jfk.wav");
4349 if repo_sample.exists() {
4350 return Ok(repo_sample);
4351 }
4352
4353 let generated = temp.join("default-alignment-smoke.wav");
4354 write_default_alignment_smoke_wav(&generated)?;
4355 Ok(generated)
4356 }
4357
4358 #[cfg(all(feature = "alignment", feature = "candle", feature = "model-bundles"))]
4359 fn resolve_alignment_smoke_bundle_candidate(path: &Path) -> Option<PathBuf> {
4360 let candidates = [
4361 path.to_path_buf(),
4362 path.join("main"),
4363 path.join("wav2vec2-base-960h"),
4364 path.join("wav2vec2-base-960h/main"),
4365 path.join("facebook--wav2vec2-base-960h"),
4366 path.join("facebook--wav2vec2-base-960h/main"),
4367 path.join("models--facebook--wav2vec2-base-960h"),
4368 ];
4369 candidates
4370 .into_iter()
4371 .find(|candidate| native_wav2vec2::resolve_wav2vec2_bundle_paths(candidate).is_ok())
4372 }
4373
4374 #[cfg(all(feature = "alignment", feature = "candle", feature = "model-bundles"))]
4375 fn default_alignment_smoke_bundle(
4376 temp: &Path,
4377 ) -> std::result::Result<(PathBuf, &'static str), Box<dyn std::error::Error>> {
4378 if let Some(bundle) = env_path("ALIGNMENT_MODEL_BUNDLE") {
4379 return Ok((bundle, "ALIGNMENT_MODEL_BUNDLE"));
4380 }
4381
4382 if let Some(model_dir) = env_path("ALIGNMENT_MODEL_DIR") {
4383 if let Some(bundle) = resolve_alignment_smoke_bundle_candidate(&model_dir) {
4384 return Ok((bundle, "ALIGNMENT_MODEL_DIR"));
4385 }
4386 return Err(format!(
4387 "ALIGNMENT_MODEL_DIR did not contain a supported wav2vec2 bundle: {}",
4388 model_dir.display()
4389 )
4390 .into());
4391 }
4392
4393 if let Some(xdg_data_home) = env_path("XDG_DATA_HOME") {
4394 let smoke_models = xdg_data_home.join("video-analysis-smoke/models");
4395 if let Some(bundle) = resolve_alignment_smoke_bundle_candidate(&smoke_models) {
4396 return Ok((bundle, "default-xdg-data-home"));
4397 }
4398 }
4399
4400 if let Some(home) = env_path("HOME") {
4401 let smoke_models = home.join(".local/share/video-analysis-smoke/models");
4402 if let Some(bundle) = resolve_alignment_smoke_bundle_candidate(&smoke_models) {
4403 return Ok((bundle, "default-home-local-share"));
4404 }
4405 }
4406
4407 let generated = temp.join("default-wav2vec2-bundle");
4408 fs::create_dir_all(&generated)?;
4409 write_tiny_wav2vec2_bundle(&generated);
4410 Ok((generated, "generated-tiny-bundle"))
4411 }
4412
4413 #[test]
4414 fn majority_overlap_assigns_word_and_segment_speaker(
4415 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4416 let mut transcript = transcript_with_words(vec![("hello", 0.1, 0.7), ("world", 0.8, 1.2)])?;
4417 let diarization = diarization_response_for_tests(vec![
4418 SpeakerSegmentPrediction {
4419 speaker: "speaker_0".to_string(),
4420 start_seconds: 0.0,
4421 end_seconds: 0.75,
4422 score: Some(0.9),
4423 },
4424 SpeakerSegmentPrediction {
4425 speaker: "speaker_1".to_string(),
4426 start_seconds: 0.75,
4427 end_seconds: 1.5,
4428 score: Some(0.8),
4429 },
4430 ]);
4431
4432 assign_speakers_from_diarization(
4433 &mut transcript,
4434 &diarization,
4435 SpeakerAssignmentPolicy::Majority,
4436 )?;
4437
4438 assert_eq!(
4439 transcript.segments[0].words[0].speaker.as_deref(),
4440 Some("speaker_0")
4441 );
4442 assert_eq!(
4443 transcript.segments[0].words[1].speaker.as_deref(),
4444 Some("speaker_1")
4445 );
4446 assert_eq!(transcript.segments[0].speaker.as_deref(), Some("speaker_0"));
4447 Ok(())
4448 }
4449
4450 #[test]
4451 fn nearest_start_policy_assigns_nearest_speaker(
4452 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4453 let mut transcript = transcript_with_words(vec![("hello", 0.72, 0.9)])?;
4454 let diarization = diarization_response_for_tests(vec![
4455 SpeakerSegmentPrediction {
4456 speaker: "speaker_0".to_string(),
4457 start_seconds: 0.0,
4458 end_seconds: 0.4,
4459 score: Some(0.9),
4460 },
4461 SpeakerSegmentPrediction {
4462 speaker: "speaker_1".to_string(),
4463 start_seconds: 0.7,
4464 end_seconds: 1.0,
4465 score: Some(0.9),
4466 },
4467 ]);
4468
4469 assign_speakers_from_diarization(
4470 &mut transcript,
4471 &diarization,
4472 SpeakerAssignmentPolicy::NearestStart,
4473 )?;
4474
4475 assert_eq!(
4476 transcript.segments[0].words[0].speaker.as_deref(),
4477 Some("speaker_1")
4478 );
4479 Ok(())
4480 }
4481
4482 #[test]
4483 fn strict_contained_policy_leaves_uncontained_words_unassigned(
4484 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4485 let mut transcript = transcript_with_words(vec![("hello", 0.2, 0.8)])?;
4486 let diarization = diarization_response_for_tests(vec![SpeakerSegmentPrediction {
4487 speaker: "speaker_0".to_string(),
4488 start_seconds: 0.3,
4489 end_seconds: 0.7,
4490 score: Some(0.9),
4491 }]);
4492
4493 assign_speakers_from_diarization(
4494 &mut transcript,
4495 &diarization,
4496 SpeakerAssignmentPolicy::StrictContained,
4497 )?;
4498
4499 assert!(transcript.segments[0].words[0].speaker.is_none());
4500 assert!(transcript.segments[0].speaker.is_none());
4501 Ok(())
4502 }
4503
4504 #[test]
4505 fn existing_segment_speaker_is_not_overwritten(
4506 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4507 let mut transcript = transcript_with_words(vec![("hello", 0.1, 0.4)])?;
4508 transcript.segments[0].speaker = Some("manual".to_string());
4509 let diarization = diarization_response_for_tests(vec![SpeakerSegmentPrediction {
4510 speaker: "speaker_0".to_string(),
4511 start_seconds: 0.0,
4512 end_seconds: 1.0,
4513 score: Some(0.9),
4514 }]);
4515
4516 assign_speakers_from_diarization(
4517 &mut transcript,
4518 &diarization,
4519 SpeakerAssignmentPolicy::Majority,
4520 )?;
4521
4522 assert_eq!(transcript.segments[0].speaker.as_deref(), Some("manual"));
4523 assert_eq!(
4524 transcript.segments[0].words[0].speaker.as_deref(),
4525 Some("speaker_0")
4526 );
4527 Ok(())
4528 }
4529
4530 #[test]
4531 fn segment_speaker_uses_majority_word_speaker_when_words_are_present(
4532 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4533 let mut transcript = transcript_with_words(vec![
4534 ("one", 0.0, 0.3),
4535 ("two", 0.3, 0.6),
4536 ("three", 0.6, 0.9),
4537 ])?;
4538 let diarization = diarization_response_for_tests(vec![
4539 SpeakerSegmentPrediction {
4540 speaker: "speaker_0".to_string(),
4541 start_seconds: 0.0,
4542 end_seconds: 0.65,
4543 score: Some(0.9),
4544 },
4545 SpeakerSegmentPrediction {
4546 speaker: "speaker_1".to_string(),
4547 start_seconds: 0.65,
4548 end_seconds: 1.0,
4549 score: Some(0.9),
4550 },
4551 ]);
4552
4553 assign_speakers_from_diarization(
4554 &mut transcript,
4555 &diarization,
4556 SpeakerAssignmentPolicy::Majority,
4557 )?;
4558
4559 assert_eq!(transcript.segments[0].speaker.as_deref(), Some("speaker_0"));
4560 Ok(())
4561 }
4562
4563 #[test]
4564 fn segment_without_words_uses_policy_fallback(
4565 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4566 let mut segment = TranscriptSegmentContract::new(0, "hello");
4567 segment.start_seconds = Some(0.5);
4568 segment.end_seconds = Some(0.8);
4569 let mut transcript =
4570 TranscriptionContract::from_segments(None, Some("en".to_string()), vec![segment])?;
4571 let diarization = diarization_response_for_tests(vec![
4572 SpeakerSegmentPrediction {
4573 speaker: "speaker_0".to_string(),
4574 start_seconds: 0.0,
4575 end_seconds: 0.2,
4576 score: Some(0.9),
4577 },
4578 SpeakerSegmentPrediction {
4579 speaker: "speaker_1".to_string(),
4580 start_seconds: 0.45,
4581 end_seconds: 1.0,
4582 score: Some(0.9),
4583 },
4584 ]);
4585
4586 assign_speakers_from_diarization(
4587 &mut transcript,
4588 &diarization,
4589 SpeakerAssignmentPolicy::StrictContained,
4590 )?;
4591
4592 assert_eq!(transcript.segments[0].speaker.as_deref(), Some("speaker_1"));
4593 Ok(())
4594 }
4595
4596 #[test]
4597 fn offset_chunk_local_segments_skips_global_timing() -> Result<()> {
4598 let mut segment = TranscriptSegmentContract::new(0, "hello");
4599 segment.start_seconds = Some(10.0);
4600 segment.end_seconds = Some(10.5);
4601 segment
4602 .attributes
4603 .insert("timing".to_string(), "global".to_string());
4604 segment.words.push(TranscriptWordContract {
4605 text: "hello".to_string(),
4606 start_seconds: Some(10.0),
4607 end_seconds: Some(10.5),
4608 confidence: None,
4609 speaker: None,
4610 attributes: BTreeMap::new(),
4611 });
4612 let mut transcript =
4613 TranscriptionContract::from_segments(None, Some("en".to_string()), vec![segment])
4614 .map_err(|error| DetectError::InvalidArgument(error.to_string()))?;
4615 let chunks = vec![SpeechActivitySegment::new(5.0, 6.0, 0.8)?];
4616
4617 offset_chunk_local_segments(&mut transcript, &chunks)?;
4618
4619 assert_eq!(transcript.segments[0].start_seconds, Some(10.0));
4620 assert_eq!(transcript.segments[0].end_seconds, Some(10.5));
4621 assert_eq!(transcript.segments[0].words[0].start_seconds, Some(10.0));
4622 assert_eq!(transcript.segments[0].words[0].end_seconds, Some(10.5));
4623 Ok(())
4624 }
4625
4626 #[test]
4627 fn offset_chunk_local_segments_offsets_local_timing() -> Result<()> {
4628 let mut segment = TranscriptSegmentContract::new(0, "hello");
4629 segment.start_seconds = Some(0.0);
4630 segment.end_seconds = Some(0.5);
4631 let mut transcript =
4632 TranscriptionContract::from_segments(None, Some("en".to_string()), vec![segment])
4633 .map_err(|error| DetectError::InvalidArgument(error.to_string()))?;
4634 let chunks = vec![SpeechActivitySegment::new(5.0, 6.0, 0.8)?];
4635
4636 offset_chunk_local_segments(&mut transcript, &chunks)?;
4637
4638 assert_eq!(transcript.segments[0].start_seconds, Some(5.0));
4639 assert_eq!(transcript.segments[0].end_seconds, Some(5.5));
4640 assert_eq!(
4641 transcript.segments[0]
4642 .attributes
4643 .get("timing")
4644 .map(String::as_str),
4645 Some("global")
4646 );
4647 Ok(())
4648 }
4649
4650 #[test]
4651 fn alignment_overwrites_projected_word_timings() -> Result<()> {
4652 let mut segment = TranscriptSegmentContract::new(0, "hello");
4653 segment.start_seconds = Some(0.0);
4654 segment.end_seconds = Some(1.0);
4655 segment.words.push(TranscriptWordContract {
4656 text: "hello".to_string(),
4657 start_seconds: Some(0.0),
4658 end_seconds: Some(0.5),
4659 confidence: None,
4660 speaker: None,
4661 attributes: BTreeMap::from([(
4662 "timing".to_string(),
4663 "whisperTimestampProjection".to_string(),
4664 )]),
4665 });
4666 let mut transcript =
4667 TranscriptionContract::from_segments(None, Some("en".to_string()), vec![segment])
4668 .map_err(|error| DetectError::InvalidArgument(error.to_string()))?;
4669
4670 apply_alignment_words(
4671 &mut transcript,
4672 &[AlignedWord {
4673 segment_index: 0,
4674 word_index: 0,
4675 text: "hello".to_string(),
4676 start_seconds: 0.1,
4677 end_seconds: 0.4,
4678 confidence: Some(0.9),
4679 }],
4680 )?;
4681
4682 let word = &transcript.segments[0].words[0];
4683 assert_eq!(word.text, "hello");
4684 assert_eq!(word.start_seconds, Some(0.1));
4685 assert_eq!(word.end_seconds, Some(0.4));
4686 assert_eq!(word.confidence, Some(0.9));
4687 assert_eq!(
4688 word.attributes.get("timing").map(String::as_str),
4689 Some("whisperTimestampProjection")
4690 );
4691 assert_eq!(transcript.segments[0].start_seconds, Some(0.1));
4692 assert_eq!(transcript.segments[0].end_seconds, Some(0.4));
4693 Ok(())
4694 }
4695
4696 #[test]
4697 fn provider_plan_reports_candle_primary_native_provider() {
4698 let plans = transcription_provider_plans();
4699 let candle = plans
4700 .iter()
4701 .find(|plan| plan.provider_id == "candle-whisper")
4702 .unwrap();
4703 assert!(candle.primary);
4704 assert!(!candle.external_runtime);
4705 assert!(plans
4706 .iter()
4707 .any(|plan| plan.provider_id == "whisperx-command" && !plan.primary));
4708 }
4709
4710 #[test]
4711 fn cuda_request_without_feature_returns_setup_error() {
4712 let mut provider = CandleWhisperTranscriber::new(CandleWhisperOptions {
4713 device: NativeDevicePreference::Cuda,
4714 ..CandleWhisperOptions::default()
4715 });
4716 let result = provider.transcribe(AsrRequest {
4717 audio: LoadedAudio {
4718 samples: vec![0.0; 16],
4719 sample_rate: 16_000,
4720 channels: 1,
4721 source: None,
4722 },
4723 chunks: vec![SpeechActivitySegment::new(0.0, 0.001, 0.0).unwrap()],
4724 task: TranscriptionTask::Transcribe,
4725 language: None,
4726 model_id: "openai/whisper-large-v3".to_string(),
4727 });
4728 let error = result.unwrap_err().to_string();
4729 assert!(error.contains("setup_error") || cfg!(feature = "cuda"));
4730 }
4731
4732 #[test]
4733 fn missing_model_bundle_returns_setup_error() {
4734 let temp = tempfile::tempdir().unwrap();
4735 let mut provider = CandleWhisperTranscriber::new(CandleWhisperOptions {
4736 model_dir: Some(temp.path().to_path_buf()),
4737 model_cache_only: true,
4738 ..CandleWhisperOptions::default()
4739 });
4740 let result = provider.transcribe(AsrRequest {
4741 audio: LoadedAudio {
4742 samples: vec![0.0; 16],
4743 sample_rate: 16_000,
4744 channels: 1,
4745 source: None,
4746 },
4747 chunks: vec![SpeechActivitySegment::new(0.0, 0.001, 0.0).unwrap()],
4748 task: TranscriptionTask::Transcribe,
4749 language: None,
4750 model_id: "openai/whisper-large-v3".to_string(),
4751 });
4752 let error = result.unwrap_err().to_string();
4753 assert!(error.contains("setup_error") || error.contains("unsupported_runtime"));
4754 assert!(error.contains("cache-only=true") || error.contains("model-bundles"));
4755 }
4756
4757 #[test]
4758 fn empty_audio_returns_invalid_request() {
4759 let mut provider = CandleWhisperTranscriber::default();
4760 let result = provider.transcribe(AsrRequest {
4761 audio: LoadedAudio {
4762 samples: Vec::new(),
4763 sample_rate: 16_000,
4764 channels: 1,
4765 source: None,
4766 },
4767 chunks: vec![SpeechActivitySegment::new(0.0, 0.001, 0.0).unwrap()],
4768 task: TranscriptionTask::Transcribe,
4769 language: None,
4770 model_id: "openai/whisper-large-v3".to_string(),
4771 });
4772 let error = result.unwrap_err().to_string();
4773 assert!(error.contains("invalid_request"));
4774 assert!(error.contains("empty audio"));
4775 }
4776
4777 #[test]
4778 fn non_finite_audio_returns_invalid_request() {
4779 let mut provider = CandleWhisperTranscriber::default();
4780 let result = provider.transcribe(AsrRequest {
4781 audio: LoadedAudio {
4782 samples: vec![0.0, f32::NAN],
4783 sample_rate: 16_000,
4784 channels: 1,
4785 source: None,
4786 },
4787 chunks: vec![SpeechActivitySegment::new(0.0, 0.001, 0.0).unwrap()],
4788 task: TranscriptionTask::Transcribe,
4789 language: None,
4790 model_id: "openai/whisper-large-v3".to_string(),
4791 });
4792 let error = result.unwrap_err().to_string();
4793 assert!(error.contains("invalid_request"));
4794 assert!(error.contains("finite"));
4795 }
4796
4797 #[cfg(not(feature = "audio-io"))]
4798 #[test]
4799 fn path_non_wav_returns_unsupported_runtime() {
4800 let result = LoadedAudio::mono_16khz_from_source(&TranscriptionSource::Path {
4801 path: PathBuf::from("clip.mp4"),
4802 });
4803 let error = result.unwrap_err().to_string();
4804 assert!(error.contains("unsupported_runtime"));
4805 assert!(error.contains("WAV"));
4806 }
4807
4808 #[test]
4809 fn wav_path_decodes_to_mono_16khz() -> std::result::Result<(), Box<dyn std::error::Error>> {
4810 let temp = tempfile::tempdir()?;
4811 let path = temp.path().join("stereo-8khz.wav");
4812 let spec = hound::WavSpec {
4813 channels: 2,
4814 sample_rate: 8_000,
4815 bits_per_sample: 16,
4816 sample_format: hound::SampleFormat::Int,
4817 };
4818 let mut writer = hound::WavWriter::create(&path, spec)?;
4819 for _ in 0..8_000 {
4820 writer.write_sample::<i16>(16_384)?;
4821 writer.write_sample::<i16>(0)?;
4822 }
4823 writer.finalize()?;
4824
4825 let audio = LoadedAudio::mono_16khz_from_source(&TranscriptionSource::Path { path })?;
4826 assert_eq!(audio.sample_rate, 16_000);
4827 assert_eq!(audio.channels, 1);
4828 assert_eq!(audio.samples.len(), 16_000);
4829 assert!(audio.samples.iter().all(|sample| sample.is_finite()));
4830 assert!(audio.samples.iter().any(|sample| *sample > 0.20));
4831 Ok(())
4832 }
4833
4834 #[test]
4835 fn decode_diagnostics_report_direct_samples_route() {
4836 let (audio, diagnostics) =
4837 native_audio::mono_16khz_from_source_with_diagnostics(&TranscriptionSource::Samples {
4838 samples: vec![0.0, 1.0],
4839 sample_rate: 8_000,
4840 channels: 1,
4841 source: Some("inline".to_string()),
4842 })
4843 .unwrap();
4844
4845 assert_eq!(diagnostics.decode_route, "direct-samples");
4846 assert_eq!(diagnostics.source_path_extension, None);
4847 assert_eq!(diagnostics.input_sample_rate, Some(8_000));
4848 assert_eq!(diagnostics.output_sample_rate, 16_000);
4849 assert_eq!(diagnostics.output_channels, 1);
4850 assert_eq!(audio.sample_rate, 16_000);
4851 assert_eq!(audio.channels, 1);
4852 }
4853
4854 #[test]
4855 fn decode_diagnostics_report_wav_route() -> std::result::Result<(), Box<dyn std::error::Error>>
4856 {
4857 let temp = tempfile::tempdir()?;
4858 let path = temp.path().join("diagnostic.wav");
4859 let spec = hound::WavSpec {
4860 channels: 1,
4861 sample_rate: 8_000,
4862 bits_per_sample: 16,
4863 sample_format: hound::SampleFormat::Int,
4864 };
4865 let mut writer = hound::WavWriter::create(&path, spec)?;
4866 writer.write_sample::<i16>(16_384)?;
4867 writer.finalize()?;
4868
4869 let (_audio, diagnostics) =
4870 native_audio::mono_16khz_from_source_with_diagnostics(&TranscriptionSource::Path {
4871 path,
4872 })?;
4873
4874 assert_eq!(diagnostics.decode_route, "native-wav-reader");
4875 assert_eq!(diagnostics.source_path_extension.as_deref(), Some("wav"));
4876 assert_eq!(diagnostics.input_sample_rate, Some(8_000));
4877 assert_eq!(diagnostics.output_sample_rate, 16_000);
4878 assert_eq!(diagnostics.output_channels, 1);
4879 Ok(())
4880 }
4881
4882 #[test]
4883 fn wav_path_still_uses_native_reader_without_audio_io(
4884 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
4885 let temp = tempfile::tempdir()?;
4886 let path = temp.path().join("native-reader.wav");
4887 let spec = hound::WavSpec {
4888 channels: 1,
4889 sample_rate: 16_000,
4890 bits_per_sample: 16,
4891 sample_format: hound::SampleFormat::Int,
4892 };
4893 let mut writer = hound::WavWriter::create(&path, spec)?;
4894 writer.write_sample::<i16>(16_384)?;
4895 writer.write_sample::<i16>(-16_384)?;
4896 writer.finalize()?;
4897
4898 let audio =
4899 LoadedAudio::mono_16khz_from_source(&TranscriptionSource::Path { path: path.clone() })?;
4900
4901 assert_eq!(audio.sample_rate, 16_000);
4902 assert_eq!(audio.channels, 1);
4903 assert_eq!(
4904 audio.source.as_deref(),
4905 Some(path.to_string_lossy().as_ref())
4906 );
4907 assert_eq!(audio.samples.len(), 2);
4908 assert!(audio.samples[0] > 0.49);
4909 assert!(audio.samples[1] < -0.49);
4910 Ok(())
4911 }
4912
4913 #[cfg(feature = "audio-io")]
4914 #[test]
4915 #[ignore = "requires RUN_NATIVE_MEDIA_DECODE_TESTS=1 and a local FFmpeg-decodable media file"]
4916 fn native_media_decode_when_requested() -> std::result::Result<(), Box<dyn std::error::Error>> {
4917 if std::env::var("RUN_NATIVE_MEDIA_DECODE_TESTS")
4918 .ok()
4919 .as_deref()
4920 != Some("1")
4921 {
4922 eprintln!("set RUN_NATIVE_MEDIA_DECODE_TESTS=1 to run native media decode smoke");
4923 return Ok(());
4924 }
4925 let path = std::env::var_os("TRANSCRIPTION_MEDIA_PATH")
4926 .map(PathBuf::from)
4927 .map(resolve_smoke_path)
4928 .ok_or("TRANSCRIPTION_MEDIA_PATH must point to a local media file")?;
4929
4930 let (audio, diagnostics) =
4931 native_audio::mono_16khz_from_source_with_diagnostics(&TranscriptionSource::Path {
4932 path: path.clone(),
4933 })?;
4934
4935 assert_eq!(audio.sample_rate, 16_000);
4936 assert_eq!(audio.channels, 1);
4937 assert!(!audio.samples.is_empty());
4938 assert!(audio.samples.iter().all(|sample| sample.is_finite()));
4939 assert_eq!(diagnostics.decode_route, "audio-io-media-decode");
4940 assert_eq!(diagnostics.output_sample_rate, 16_000);
4941 assert_eq!(diagnostics.output_channels, 1);
4942 assert!(diagnostics.input_sample_rate.is_some());
4943 assert_eq!(
4944 audio.source.as_deref(),
4945 Some(path.to_string_lossy().as_ref())
4946 );
4947 Ok(())
4948 }
4949
4950 #[cfg(feature = "audio-io")]
4951 fn resolve_smoke_path(path: PathBuf) -> PathBuf {
4952 if path.is_absolute() || path.exists() {
4953 return path;
4954 }
4955 let manifest_dir = Path::new(env!("CARGO_MANIFEST_DIR"));
4956 let Some(workspace_root) = manifest_dir.ancestors().nth(3) else {
4957 return path;
4958 };
4959 let workspace_path = workspace_root.join(&path);
4960 if workspace_path.exists() {
4961 workspace_path
4962 } else {
4963 path
4964 }
4965 }
4966
4967 #[test]
4968 fn vad_splits_deterministic_synthetic_speech() {
4969 let request = sample_request();
4970 let audio = LoadedAudio::mono_16khz_from_source(&request.source).unwrap();
4971 let mut vad = EnergyVadTranscriptionProvider;
4972 let response = vad
4973 .detect_speech(VadRequest {
4974 audio,
4975 options: request.vad,
4976 })
4977 .unwrap();
4978 assert_eq!(response.segments.len(), 1);
4979 assert!(response.segments[0].start_seconds < 0.07);
4980 assert!(response.segments[0].end_seconds > 0.30);
4981 }
4982
4983 #[test]
4984 fn mock_pipeline_normalizes_offsets_alignment_and_diarization() {
4985 let mut request = sample_request();
4986 request.alignment.enabled = true;
4987 request.diarization.enabled = true;
4988 let mut vad = EnergyVadTranscriptionProvider;
4989 let mut asr = MockAsrProvider;
4990 let mut aligner = MockAlignmentProvider;
4991 let mut diarizer = MockDiarizationProvider;
4992 let response = run_transcription_pipeline(
4993 request,
4994 &mut vad,
4995 &mut asr,
4996 Some(&mut aligner),
4997 Some(&mut diarizer),
4998 )
4999 .unwrap();
5000 assert!(response.accepted);
5001 assert_eq!(response.provider, "candle-whisper");
5002 assert_eq!(response.alignment.as_ref().unwrap().word_count, 1);
5003 assert_eq!(
5004 response.transcript.segments[0].words[0].speaker.as_deref(),
5005 Some("SPEAKER_00")
5006 );
5007 assert_eq!(
5008 response.transcript.segments[0].speaker.as_deref(),
5009 Some("SPEAKER_00")
5010 );
5011 }
5012
5013 #[test]
5014 fn native_pipeline_reports_diarization_diagnostics() {
5015 let mut request = sample_request();
5016 request.diarization = DiarizationOptions {
5017 enabled: true,
5018 speaker: SpeakerDiarizationOptions {
5019 min_speakers: Some(2),
5020 max_speakers: Some(3),
5021 ..SpeakerDiarizationOptions::default()
5022 },
5023 };
5024 let mut vad = EnergyVadTranscriptionProvider;
5025 let mut asr = MockAsrProvider;
5026 let mut diarizer = MockDiarizationProvider;
5027
5028 let response =
5029 run_transcription_pipeline(request, &mut vad, &mut asr, None, Some(&mut diarizer))
5030 .unwrap();
5031
5032 assert!(response
5033 .diagnostics
5034 .iter()
5035 .any(|item| item == "diarizationProvider=mock-diarization"));
5036 assert!(response
5037 .diagnostics
5038 .iter()
5039 .any(|item| item == "diarizationSegmentCount=1"));
5040 assert!(response
5041 .diagnostics
5042 .iter()
5043 .any(|item| item == "diarizationSpeakerCount=1"));
5044 assert!(response
5045 .diagnostics
5046 .iter()
5047 .any(|item| item == "diarizationMinSpeakers=2"));
5048 assert!(response
5049 .diagnostics
5050 .iter()
5051 .any(|item| item == "diarizationSpeakerBoundsApplied=true"));
5052 assert!(response
5053 .diagnostics
5054 .iter()
5055 .any(|item| item == "diarizationSpeakerCountBelowRequestedMin=1/2"));
5056 }
5057
5058 #[cfg(feature = "diarization")]
5059 #[test]
5060 fn native_pipeline_reports_onnx_diarization_diagnostics() {
5061 let mut request = sample_request();
5062 request.diarization = DiarizationOptions {
5063 enabled: true,
5064 speaker: SpeakerDiarizationOptions {
5065 speaker_embedding_model_bundle: Some(PathBuf::from("speaker-model")),
5066 speaker_embedding_dimension: Some(2),
5067 ..SpeakerDiarizationOptions::default()
5068 },
5069 ..DiarizationOptions::default()
5070 };
5071 let mut vad = EnergyVadTranscriptionProvider;
5072 let mut asr = MockAsrProvider;
5073 let mut diarizer = MockOnnxDiarizationProvider;
5074
5075 let response =
5076 run_transcription_pipeline(request, &mut vad, &mut asr, None, Some(&mut diarizer))
5077 .unwrap();
5078
5079 assert!(response
5080 .diagnostics
5081 .iter()
5082 .any(|item| item == "diarizationRuntime=onnx"));
5083 assert!(response
5084 .diagnostics
5085 .iter()
5086 .any(|item| item == "speakerEmbeddingProvider=onnx"));
5087 assert!(response
5088 .diagnostics
5089 .iter()
5090 .any(|item| item == "speakerEmbeddingDimension=2"));
5091 assert!(response
5092 .diagnostics
5093 .iter()
5094 .any(|item| item == "diarizationBaseline=false"));
5095 assert!(!response
5096 .diagnostics
5097 .iter()
5098 .any(|item| item == "diarizationBaseline=heuristic-native"));
5099 assert_eq!(
5100 response.transcript.segments[0].speaker.as_deref(),
5101 Some("SPEAKER_ONNX")
5102 );
5103 }
5104
5105 #[cfg(feature = "diarization")]
5106 #[test]
5107 fn native_onnx_diarization_missing_bundle_returns_setup_error() {
5108 let mut request = sample_request();
5109 request.diarization = DiarizationOptions {
5110 enabled: true,
5111 speaker: SpeakerDiarizationOptions {
5112 speaker_embedding_model_bundle: Some(PathBuf::from(
5113 "/definitely/missing/onnx-speaker-model",
5114 )),
5115 speaker_embedding_dimension: Some(2),
5116 ..SpeakerDiarizationOptions::default()
5117 },
5118 ..DiarizationOptions::default()
5119 };
5120 let mut vad = EnergyVadTranscriptionProvider;
5121 let mut asr = MockAsrProvider;
5122 let mut diarizer = NativeSpeakerDiarizationProvider;
5123
5124 let error =
5125 run_native_transcription_pipeline(request, &mut vad, &mut asr, Some(&mut diarizer))
5126 .unwrap_err()
5127 .to_string();
5128
5129 assert!(error.contains("setup_error"));
5130 assert!(error.contains("ONNX speaker embedding"));
5131 }
5132
5133 #[cfg(all(feature = "diarization", feature = "onnx"))]
5134 #[test]
5135 #[ignore = "requires RUN_NATIVE_SPEAKER_MODEL_TESTS=1 and caller-owned local ONNX speaker model/audio"]
5136 fn native_onnx_diarization_smoke_when_requested(
5137 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5138 if std::env::var("RUN_NATIVE_SPEAKER_MODEL_TESTS").as_deref() != Ok("1") {
5139 eprintln!(
5140 "skipping native ONNX diarization smoke; set RUN_NATIVE_SPEAKER_MODEL_TESTS=1"
5141 );
5142 return Ok(());
5143 }
5144
5145 let bundle_path = std::env::var_os("SPEAKER_EMBEDDING_MODEL_BUNDLE")
5146 .map(PathBuf::from)
5147 .ok_or("SPEAKER_EMBEDDING_MODEL_BUNDLE is required")?;
5148 let audio_path = std::env::var_os("DIARIZATION_AUDIO_PATH")
5149 .map(PathBuf::from)
5150 .ok_or("DIARIZATION_AUDIO_PATH is required")?;
5151 let embedding_dimension = std::env::var("SPEAKER_EMBEDDING_DIMENSION")
5152 .ok()
5153 .map(|value| value.parse::<usize>())
5154 .transpose()?
5155 .unwrap_or(192);
5156 let model_file =
5157 optional_smoke_env_value(std::env::var("SPEAKER_EMBEDDING_MODEL_FILE").ok());
5158 let input_name =
5159 optional_smoke_env_value(std::env::var("SPEAKER_EMBEDDING_INPUT_NAME").ok());
5160 let output_name =
5161 optional_smoke_env_value(std::env::var("SPEAKER_EMBEDDING_OUTPUT_NAME").ok());
5162 let model_path = resolve_onnx_smoke_model_path(&bundle_path, model_file.as_deref())?;
5163 eprintln!("speakerEmbeddingResolvedModelPath={}", model_path.display());
5164 eprintln!("speakerEmbeddingExpectedDimension={embedding_dimension}");
5165 eprintln!(
5166 "speakerEmbeddingConfiguredInputName={}",
5167 input_name.as_deref().unwrap_or("<auto>")
5168 );
5169 eprintln!(
5170 "speakerEmbeddingConfiguredOutputName={}",
5171 output_name.as_deref().unwrap_or("<auto>")
5172 );
5173 let static_metadata = runtime_onnx::inspect_model_metadata(&model_path)?;
5174 eprintln!("onnxStaticMetadata=ok");
5175 for diagnostic in runtime_onnx::inspect_model_graph_diagnostics(&model_path)? {
5176 eprintln!("{diagnostic}");
5177 }
5178 eprintln!(
5179 "onnxLoadMode={}",
5180 std::env::var("ONNX_RUNTIME_LOAD_MODE").unwrap_or_else(|_| "file".to_string())
5181 );
5182 eprintln!(
5183 "onnxRuntimeDylib={}",
5184 if std::env::var_os("ORT_DYLIB_PATH").is_some() {
5185 "set"
5186 } else {
5187 "unset"
5188 }
5189 );
5190 if let Some(input) = static_metadata.inputs.first() {
5191 eprintln!("speakerEmbeddingStaticInputName={}", input.name);
5192 eprintln!(
5193 "speakerEmbeddingStaticInputDimensions={}",
5194 format_onnx_smoke_dimensions(&input.dimensions)
5195 );
5196 }
5197 if let Some(output) = static_metadata.outputs.first() {
5198 eprintln!("speakerEmbeddingStaticOutputName={}", output.name);
5199 eprintln!(
5200 "speakerEmbeddingStaticOutputDimensions={}",
5201 format_onnx_smoke_dimensions(&output.dimensions)
5202 );
5203 }
5204 eprintln!(
5205 "onnxSessionOptions=cpu-single-threaded,no-memory-pattern,graph-optimization-disabled"
5206 );
5207 let samples = local_16khz_wav_samples(&audio_path)?;
5208 let duration_seconds = samples.len() as f64 / 16_000.0;
5209 let midpoint = (duration_seconds / 2.0).clamp(1.0 / 16_000.0, duration_seconds);
5210 let first_end = midpoint.min(duration_seconds);
5211 let second_start = first_end;
5212 let second_end = duration_seconds;
5213
5214 let mut request = TranscriptionPipelineRequest {
5215 source: TranscriptionSource::Samples {
5216 samples,
5217 sample_rate: 16_000,
5218 channels: 1,
5219 source: Some(audio_path.to_string_lossy().into_owned()),
5220 },
5221 provider: TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions::default()),
5222 vad: VadOptions::default(),
5223 alignment: AlignmentOptions::default(),
5224 diarization: DiarizationOptions {
5225 enabled: true,
5226 speaker: SpeakerDiarizationOptions {
5227 speaker_embedding_model_bundle: Some(bundle_path),
5228 speaker_embedding_model_file: model_file,
5229 speaker_embedding_input_name: input_name,
5230 speaker_embedding_output_name: output_name,
5231 speaker_embedding_dimension: Some(embedding_dimension),
5232 speaker_embedding_sample_rate: Some(16_000),
5233 assignment_policy: SpeakerAssignmentPolicy::StrictContained,
5234 ..SpeakerDiarizationOptions::default()
5235 },
5236 ..DiarizationOptions::default()
5237 },
5238 output: TranscriptionOutputOptions::default(),
5239 };
5240 if second_end > second_start {
5241 request.vad.enabled = true;
5242 }
5243
5244 let segments = if second_end > second_start {
5245 vec![
5246 SpeechActivitySegment::new(0.0, first_end, 1.0)?,
5247 SpeechActivitySegment::new(second_start, second_end, 1.0)?,
5248 ]
5249 } else {
5250 vec![SpeechActivitySegment::new(0.0, first_end, 1.0)?]
5251 };
5252 let mut vad = FixedVadProvider { segments };
5253 let mut asr = MockAsrProvider;
5254 let mut diarizer = NativeSpeakerDiarizationProvider;
5255
5256 let response =
5257 run_native_transcription_pipeline(request, &mut vad, &mut asr, Some(&mut diarizer))?;
5258 eprintln!("{}", response.diagnostics.join("\n"));
5259
5260 assert!(response.accepted);
5261 assert!(response
5262 .diagnostics
5263 .iter()
5264 .any(|item| item == "diarizationRuntime=onnx"));
5265 assert!(response
5266 .diagnostics
5267 .iter()
5268 .any(|item| item == "speakerEmbeddingProvider=onnx"));
5269 assert!(response
5270 .diagnostics
5271 .iter()
5272 .any(|item| item == &format!("speakerEmbeddingDimension={embedding_dimension}")));
5273 assert!(response
5274 .diagnostics
5275 .iter()
5276 .any(|item| item == "diarizationBaseline=false"));
5277 assert!(response.transcript.segments.iter().all(|segment| segment
5278 .speaker
5279 .as_deref()
5280 .is_some_and(|speaker| { !speaker.trim().is_empty() })));
5281 normalize_transcription_contract(response.transcript.clone())
5282 .map_err(|error| format!("transcript speaker assignment must validate: {error}"))?;
5283 Ok(())
5284 }
5285
5286 #[cfg(all(feature = "diarization", feature = "onnx"))]
5287 fn optional_smoke_env_value(value: Option<String>) -> Option<String> {
5288 value.filter(|value| !value.trim().is_empty())
5289 }
5290
5291 #[cfg(all(feature = "diarization", feature = "onnx"))]
5292 fn resolve_onnx_smoke_model_path(
5293 bundle_path: &Path,
5294 model_file: Option<&str>,
5295 ) -> std::result::Result<PathBuf, Box<dyn std::error::Error>> {
5296 let model_file = model_file.unwrap_or("model.onnx");
5297 if bundle_path.is_file() {
5298 return Ok(bundle_path.to_path_buf());
5299 }
5300 #[cfg(feature = "model-bundles")]
5301 {
5302 let manifest_path = bundle_path.join("manifest.json");
5303 if manifest_path.is_file() {
5304 let bundle = model_runtime::ModelBundle::load(&manifest_path)?;
5305 for file in bundle.manifest.files.values() {
5306 if file.remote_path == model_file
5307 || file.remote_path.ends_with(model_file)
5308 || file.local_path.ends_with(model_file)
5309 {
5310 if let Some(path) = bundle.file_path(&file.remote_path) {
5311 return Ok(path);
5312 }
5313 }
5314 }
5315 }
5316 }
5317 Ok(bundle_path.join(model_file))
5318 }
5319
5320 #[cfg(all(feature = "diarization", feature = "onnx"))]
5321 fn format_onnx_smoke_dimensions(dimensions: &[runtime_onnx::OnnxDimension]) -> String {
5322 let values = dimensions
5323 .iter()
5324 .map(|dimension| match dimension {
5325 runtime_onnx::OnnxDimension::Fixed(value) => value.to_string(),
5326 runtime_onnx::OnnxDimension::Symbolic(value) => value.clone(),
5327 runtime_onnx::OnnxDimension::Unknown => "unknown".to_string(),
5328 })
5329 .collect::<Vec<_>>()
5330 .join(",");
5331 format!("[{values}]")
5332 }
5333
5334 #[test]
5335 fn native_pipeline_diarization_invalid_bounds_errors_before_provider() {
5336 let mut request = sample_request();
5337 request.diarization = DiarizationOptions {
5338 enabled: true,
5339 speaker: SpeakerDiarizationOptions {
5340 min_speakers: Some(3),
5341 max_speakers: Some(2),
5342 ..SpeakerDiarizationOptions::default()
5343 },
5344 };
5345 let mut vad = EnergyVadTranscriptionProvider;
5346 let mut asr = MockAsrProvider;
5347 let mut diarizer = PanickingDiarizationProvider { called: false };
5348
5349 let error =
5350 run_transcription_pipeline(request, &mut vad, &mut asr, None, Some(&mut diarizer))
5351 .unwrap_err()
5352 .to_string();
5353
5354 assert!(error.contains("invalid_request"));
5355 assert!(!diarizer.called);
5356 }
5357
5358 #[test]
5359 #[cfg(not(feature = "diarization"))]
5360 fn native_diarization_without_feature_still_reports_no_provider() {
5361 let mut request = sample_request();
5362 request.diarization.enabled = true;
5363 let mut vad = EnergyVadTranscriptionProvider;
5364 let mut asr = MockAsrProvider;
5365
5366 let error = run_transcription_pipeline(request, &mut vad, &mut asr, None, None)
5367 .unwrap_err()
5368 .to_string();
5369
5370 assert!(error.contains("setup_error"));
5371 assert!(error.contains("no diarization provider is available"));
5372 }
5373
5374 #[test]
5375 #[cfg(feature = "diarization")]
5376 fn native_speaker_diarization_provider_uses_transcript_spans_when_available(
5377 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5378 let samples = (0..32_000)
5379 .map(|index| if index % 80 < 40 { 0.2 } else { -0.2 })
5380 .collect::<Vec<_>>();
5381 let audio = LoadedAudio {
5382 samples,
5383 sample_rate: 16_000,
5384 channels: 1,
5385 source: Some("synthetic".to_string()),
5386 };
5387 let transcript = transcript_with_words(vec![("hello", 0.20, 0.50), ("world", 1.00, 1.40)])?;
5388 let options = DiarizationOptions {
5389 enabled: true,
5390 speaker: SpeakerDiarizationOptions {
5391 model_id: "requested-native-speakers".to_string(),
5392 ..SpeakerDiarizationOptions::default()
5393 },
5394 ..DiarizationOptions::default()
5395 };
5396 let mut provider = NativeSpeakerDiarizationProvider;
5397
5398 let response = provider.diarize(audio, &transcript, &options)?;
5399
5400 assert_eq!(response.model_id, "requested-native-speakers");
5401 assert_eq!(
5402 response.runtime,
5403 audio_analysis_speakers::AudioRuntime::Heuristic
5404 );
5405 assert_eq!(response.segments.len(), 2);
5406 assert!((response.segments[0].start_seconds - 0.20).abs() < 0.001);
5407 assert!((response.segments[0].end_seconds - 0.50).abs() < 0.001);
5408 assert!((response.segments[1].start_seconds - 1.00).abs() < 0.001);
5409 assert!((response.segments[1].end_seconds - 1.40).abs() < 0.001);
5410 assert!(response
5411 .segments
5412 .iter()
5413 .all(|segment| segment.speaker.starts_with("speaker_")));
5414 Ok(())
5415 }
5416
5417 #[test]
5418 #[cfg(feature = "diarization")]
5419 fn native_speaker_diarization_provider_applies_exact_speaker_bounds(
5420 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5421 let audio = two_profile_loaded_audio();
5422 let transcript = transcript_with_words(vec![("hello", 0.20, 0.50), ("world", 1.00, 1.40)])?;
5423 let options = DiarizationOptions {
5424 enabled: true,
5425 speaker: SpeakerDiarizationOptions {
5426 min_speakers: Some(2),
5427 max_speakers: Some(2),
5428 ..SpeakerDiarizationOptions::default()
5429 },
5430 ..DiarizationOptions::default()
5431 };
5432 let mut provider = NativeSpeakerDiarizationProvider;
5433
5434 let response = provider.diarize(audio, &transcript, &options)?;
5435 let speakers = response
5436 .segments
5437 .iter()
5438 .map(|segment| segment.speaker.clone())
5439 .collect::<std::collections::BTreeSet<_>>();
5440
5441 assert_eq!(speakers.len(), 2, "{:?}", response.segments);
5442 assert!(speakers.contains("speaker_0"));
5443 assert!(speakers.contains("speaker_1"));
5444 Ok(())
5445 }
5446
5447 #[test]
5448 #[cfg(feature = "diarization")]
5449 fn native_speaker_diarization_provider_applies_max_one_bound(
5450 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5451 let audio = two_profile_loaded_audio();
5452 let transcript = transcript_with_words(vec![("hello", 0.20, 0.50), ("world", 1.00, 1.40)])?;
5453 let options = DiarizationOptions {
5454 enabled: true,
5455 speaker: SpeakerDiarizationOptions {
5456 max_speakers: Some(1),
5457 ..SpeakerDiarizationOptions::default()
5458 },
5459 ..DiarizationOptions::default()
5460 };
5461 let mut provider = NativeSpeakerDiarizationProvider;
5462
5463 let response = provider.diarize(audio, &transcript, &options)?;
5464 let speakers = response
5465 .segments
5466 .iter()
5467 .map(|segment| segment.speaker.clone())
5468 .collect::<std::collections::BTreeSet<_>>();
5469
5470 assert_eq!(speakers.len(), 1, "{:?}", response.segments);
5471 assert!(speakers.contains("speaker_0"));
5472 Ok(())
5473 }
5474
5475 #[test]
5476 #[cfg(feature = "diarization")]
5477 fn native_speaker_diarization_provider_falls_back_without_transcript_timing(
5478 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5479 let samples = (0..16_000)
5480 .map(|index| {
5481 if (1_000..8_000).contains(&index) {
5482 0.2
5483 } else {
5484 0.0
5485 }
5486 })
5487 .collect::<Vec<_>>();
5488 let audio = LoadedAudio {
5489 samples,
5490 sample_rate: 16_000,
5491 channels: 1,
5492 source: Some("synthetic".to_string()),
5493 };
5494 let segment = TranscriptSegmentContract::new(0, "hello without timing");
5495 let transcript =
5496 TranscriptionContract::from_segments(None, Some("en".to_string()), vec![segment])?;
5497 let options = DiarizationOptions {
5498 enabled: true,
5499 speaker: SpeakerDiarizationOptions {
5500 model_id: "fallback-native-speakers".to_string(),
5501 ..SpeakerDiarizationOptions::default()
5502 },
5503 ..DiarizationOptions::default()
5504 };
5505 let mut provider = NativeSpeakerDiarizationProvider;
5506
5507 let response = provider.diarize(audio, &transcript, &options)?;
5508
5509 assert_eq!(response.model_id, "fallback-native-speakers");
5510 assert_eq!(response.operation, "audio.speakers.diarize");
5511 assert_eq!(
5512 response.runtime,
5513 audio_analysis_speakers::AudioRuntime::Heuristic
5514 );
5515 assert!(!response.segments.is_empty());
5516 Ok(())
5517 }
5518
5519 #[test]
5520 fn alignment_options_default_device_is_cpu() {
5521 assert_eq!(
5522 AlignmentOptions::default().device,
5523 NativeDevicePreference::Cpu
5524 );
5525 }
5526
5527 #[test]
5528 fn alignment_options_deserializes_cuda_device() {
5529 let options: AlignmentOptions =
5530 serde_json::from_str(r#"{"device":"cuda"}"#).expect("alignment options should parse");
5531
5532 assert_eq!(options.device, NativeDevicePreference::Cuda);
5533 }
5534
5535 #[test]
5536 #[cfg(feature = "alignment")]
5537 fn native_pipeline_supplies_alignment_provider_when_enabled(
5538 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5539 let temp = tempfile::tempdir()?;
5540 write_tiny_wav2vec2_bundle(temp.path());
5541 let mut request = sample_request();
5542 request.alignment = AlignmentOptions {
5543 enabled: true,
5544 model_bundle: Some(temp.path().to_path_buf()),
5545 ..AlignmentOptions::default()
5546 };
5547 let mut vad = EnergyVadTranscriptionProvider;
5548 let mut asr = MockAsrProvider;
5549
5550 let response = run_native_transcription_pipeline(request, &mut vad, &mut asr, None)?;
5551
5552 let alignment = response.alignment.as_ref().unwrap();
5553 assert_eq!(alignment.provider, "ctc-forced-aligner");
5554 assert_eq!(alignment.model_id, default_alignment_model());
5555 assert_eq!(alignment.word_count, 1);
5556 let word = &response.transcript.segments[0].words[0];
5557 assert_eq!(word.text, "hello");
5558 assert!(word.start_seconds.is_some());
5559 assert!(word.end_seconds.is_some());
5560 assert!(word
5561 .confidence
5562 .is_some_and(|confidence| confidence.is_finite()));
5563 assert!(response
5564 .diagnostics
5565 .iter()
5566 .any(|item| item == "alignmentModelSource=explicit-bundle"));
5567 assert!(response
5568 .diagnostics
5569 .iter()
5570 .any(|item| item == "alignmentDevice=cpu"));
5571 assert!(response
5572 .diagnostics
5573 .iter()
5574 .any(|item| item == "alignmentCuda=false"));
5575 response.transcript.validate_strict()?;
5576 Ok(())
5577 }
5578
5579 #[test]
5580 fn native_pipeline_leaves_alignment_absent_when_disabled() {
5581 let request = sample_request();
5582 let mut vad = EnergyVadTranscriptionProvider;
5583 let mut asr = MockAsrProvider;
5584
5585 let response =
5586 run_native_transcription_pipeline(request, &mut vad, &mut asr, None).unwrap();
5587
5588 assert!(response.alignment.is_none());
5589 assert!(response
5590 .diagnostics
5591 .iter()
5592 .all(|item| !item.to_lowercase().contains("alignment")));
5593 response
5594 .transcript
5595 .validate_strict()
5596 .expect("native pipeline response should be strictly valid");
5597 }
5598
5599 #[test]
5600 fn native_pipeline_passes_translate_task_to_asr_provider() {
5601 let mut request = sample_request();
5602 request.provider = TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions {
5603 task: TranscriptionTask::Translate,
5604 ..CandleWhisperOptions::default()
5605 });
5606 let mut vad = EnergyVadTranscriptionProvider;
5607 let mut asr = MockAsrProvider;
5608
5609 let response =
5610 run_native_transcription_pipeline(request, &mut vad, &mut asr, None).unwrap();
5611
5612 assert!(response
5613 .diagnostics
5614 .iter()
5615 .any(|item| item == "asrTask=translate"));
5616 assert!(response
5617 .diagnostics
5618 .iter()
5619 .any(|item| item == "translationRuntime=whisper-task"));
5620 assert!(response
5621 .diagnostics
5622 .iter()
5623 .any(|item| item == "translationTargetLanguage=en"));
5624 assert_eq!(response.transcript.language.as_deref(), Some("en"));
5625 assert!(response.alignment.is_none());
5626 }
5627
5628 #[test]
5629 fn native_translate_rejects_ctc_alignment() {
5630 let mut request = sample_request();
5631 request.provider = TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions {
5632 task: TranscriptionTask::Translate,
5633 ..CandleWhisperOptions::default()
5634 });
5635 request.alignment.enabled = true;
5636 let mut vad = EnergyVadTranscriptionProvider;
5637 let mut asr = MockAsrProvider;
5638
5639 let error = run_native_transcription_pipeline(request, &mut vad, &mut asr, None)
5640 .unwrap_err()
5641 .to_string();
5642
5643 assert!(error.contains("invalid_request"));
5644 assert!(error.contains("translation output cannot be wav2vec2/CTC-aligned"));
5645 }
5646
5647 #[test]
5648 fn native_translate_allows_diarization_with_segment_timings() {
5649 let mut request = sample_request();
5650 request.provider = TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions {
5651 task: TranscriptionTask::Translate,
5652 ..CandleWhisperOptions::default()
5653 });
5654 request.diarization.enabled = true;
5655 let mut vad = EnergyVadTranscriptionProvider;
5656 let mut asr = MockAsrProvider;
5657 let mut diarizer = MockDiarizationProvider;
5658
5659 let response =
5660 run_native_transcription_pipeline(request, &mut vad, &mut asr, Some(&mut diarizer))
5661 .unwrap();
5662
5663 assert!(response.diarization.is_some());
5664 assert_eq!(
5665 response.transcript.segments[0].speaker.as_deref(),
5666 Some("SPEAKER_00")
5667 );
5668 assert!(response
5669 .diagnostics
5670 .iter()
5671 .any(|item| item == "asrTask=translate"));
5672 }
5673
5674 #[test]
5675 #[cfg(feature = "alignment")]
5676 fn native_pipeline_runs_alignment_before_diarization() {
5677 let temp = tempfile::tempdir().expect("tempdir");
5678 write_tiny_wav2vec2_bundle(temp.path());
5679 let mut request = sample_request();
5680 request.alignment = AlignmentOptions {
5681 enabled: true,
5682 model_bundle: Some(temp.path().to_path_buf()),
5683 ..AlignmentOptions::default()
5684 };
5685 request.diarization.enabled = true;
5686 let mut vad = EnergyVadTranscriptionProvider;
5687 let mut asr = MockAsrProvider;
5688 let mut diarizer = AlignmentAwareDiarizationProvider {
5689 saw_aligned_words: false,
5690 };
5691
5692 let response =
5693 run_native_transcription_pipeline(request, &mut vad, &mut asr, Some(&mut diarizer))
5694 .unwrap();
5695
5696 assert!(diarizer.saw_aligned_words);
5697 assert!(response.diarization.is_some());
5698 let word = &response.transcript.segments[0].words[0];
5699 assert!(word.start_seconds.is_some());
5700 assert!(word.end_seconds.is_some());
5701 assert_eq!(word.speaker.as_deref(), Some("SPEAKER_ALIGNED"));
5702 assert_eq!(
5703 response.transcript.segments[0].speaker.as_deref(),
5704 Some("SPEAKER_ALIGNED")
5705 );
5706 }
5707
5708 #[test]
5709 #[cfg(not(feature = "alignment"))]
5710 fn native_pipeline_alignment_without_feature_reports_alignment_unsupported_runtime() {
5711 let mut request = sample_request();
5712 request.alignment.enabled = true;
5713 let mut vad = EnergyVadTranscriptionProvider;
5714 let mut asr = MockAsrProvider;
5715
5716 let error = run_native_transcription_pipeline(request, &mut vad, &mut asr, None)
5717 .unwrap_err()
5718 .to_string();
5719
5720 assert!(error.contains("unsupported_runtime"));
5721 assert!(error.contains("CTC alignment"));
5722 assert!(error.contains("alignment"));
5723 assert!(!error.contains("no alignment provider is available"));
5724 }
5725
5726 #[test]
5727 #[cfg(all(feature = "alignment", feature = "candle"))]
5728 fn native_pipeline_with_tiny_wav2vec2_bundle_runs_model_backed_alignment(
5729 ) -> std::result::Result<(), Box<dyn std::error::Error>> {
5730 let temp = tempfile::tempdir()?;
5731 write_tiny_wav2vec2_bundle(temp.path());
5732 let mut request = sample_request();
5733 request.alignment = AlignmentOptions {
5734 enabled: true,
5735 model_bundle: Some(temp.path().to_path_buf()),
5736 ..AlignmentOptions::default()
5737 };
5738 let mut vad = EnergyVadTranscriptionProvider;
5739 let mut asr = MockAsrProvider;
5740
5741 let response = run_native_transcription_pipeline(request, &mut vad, &mut asr, None)?;
5742
5743 let alignment = response.alignment.as_ref().unwrap();
5744 assert_eq!(alignment.provider, "ctc-forced-aligner");
5745 assert_eq!(alignment.word_count, 1);
5746 assert!(response
5747 .diagnostics
5748 .iter()
5749 .any(|item| item == "alignmentModelExecution=candle-wav2vec2"));
5750 let word = &response.transcript.segments[0].words[0];
5751 assert_eq!(word.text, "hello");
5752 assert!(word.start_seconds.is_some());
5753 assert!(word.end_seconds.is_some());
5754 assert!(word.confidence.is_some());
5755 response.transcript.validate_strict()?;
5756 Ok(())
5757 }
5758
5759 #[test]
5760 fn missing_command_returns_setup_error() {
5761 let result = transcribe(TranscriptionPipelineRequest {
5762 source: TranscriptionSource::Path {
5763 path: PathBuf::from("missing.wav"),
5764 },
5765 provider: TranscriptionProviderSelection::ExternalWhisperX(WhisperXCommandOptions {
5766 command: PathBuf::from("definitely-missing-whisperx-command"),
5767 ..WhisperXCommandOptions::default()
5768 }),
5769 vad: VadOptions::default(),
5770 alignment: AlignmentOptions::default(),
5771 diarization: DiarizationOptions::default(),
5772 output: TranscriptionOutputOptions::default(),
5773 });
5774
5775 let error = result.unwrap_err().to_string();
5776 assert!(error.contains("setup_error"));
5777 assert!(error.contains("not found"));
5778 }
5779
5780 #[test]
5781 fn whisperx_json_bytes_prefers_current_source_stem() {
5782 let temp = tempfile::tempdir().expect("tempdir");
5783 fs::write(temp.path().join("a.json"), br#"{"text":"first"}"#).expect("a json");
5784 fs::write(temp.path().join("b.json"), br#"{"text":"second"}"#).expect("b json");
5785
5786 let (path, bytes) =
5787 whisperx_json_bytes(Path::new("audio/b.wav"), temp.path(), b"{}").expect("json");
5788
5789 assert_eq!(path, Some(temp.path().join("b.json")));
5790 assert_eq!(bytes, br#"{"text":"second"}"#);
5791 }
5792
5793 #[test]
5794 fn reusable_candle_provider_reports_unsupported_without_candle_feature() {
5795 #[cfg(not(feature = "candle"))]
5796 {
5797 let mut provider = ReusableCandleWhisperTranscriber::new(CandleWhisperOptions {
5798 model_bundle: Some(PathBuf::from("bundle")),
5799 ..CandleWhisperOptions::default()
5800 });
5801 let error = provider
5802 .transcribe(AsrRequest {
5803 audio: LoadedAudio {
5804 samples: vec![0.0; 16],
5805 sample_rate: 16_000,
5806 channels: 1,
5807 source: None,
5808 },
5809 chunks: vec![SpeechActivitySegment::new(0.0, 0.001, 1.0).unwrap()],
5810 task: TranscriptionTask::Transcribe,
5811 language: Some("en".to_string()),
5812 model_id: "tiny.en".to_string(),
5813 })
5814 .unwrap_err()
5815 .to_string();
5816
5817 assert!(
5818 error.contains("unsupported_runtime") || error.contains("setup_error"),
5819 "{error}"
5820 );
5821 assert!(
5822 error.contains("candle") || error.contains("bundle"),
5823 "{error}"
5824 );
5825 }
5826 }
5827
5828 #[test]
5829 fn diarization_requires_token_before_spawn() {
5830 let result = transcribe(TranscriptionPipelineRequest {
5831 source: TranscriptionSource::Path {
5832 path: PathBuf::from("missing.wav"),
5833 },
5834 provider: TranscriptionProviderSelection::ExternalWhisperX(WhisperXCommandOptions {
5835 command: PathBuf::from("definitely-missing-whisperx-command"),
5836 diarize: true,
5837 hf_token_env: Some("VIDEO_ANALYSIS_TEST_MISSING_HF_TOKEN".to_string()),
5838 ..WhisperXCommandOptions::default()
5839 }),
5840 vad: VadOptions::default(),
5841 alignment: AlignmentOptions::default(),
5842 diarization: DiarizationOptions::default(),
5843 output: TranscriptionOutputOptions::default(),
5844 });
5845
5846 let error = result.unwrap_err().to_string();
5847 assert!(error.contains("diarization requires"));
5848 assert!(!error.contains("not found"));
5849 }
5850
5851 #[test]
5852 fn mock_command_output_round_trips() -> std::result::Result<(), Box<dyn std::error::Error>> {
5853 let temp = tempfile::tempdir()?;
5854 let command = temp.path().join("mock-whisperx.sh");
5855 let output_dir = temp.path().join("out");
5856 fs::write(
5857 &command,
5858 format!(
5859 "#!/usr/bin/env bash\nmkdir -p \"{}\"\nprintf 'Transcript: [0.29 --> 1.47] hello\\n'\ncat > \"{}/sample.json\" <<'JSON'\n{{\"segments\":[{{\"start\":0.0,\"end\":1.0,\"text\":\" hello \",\"speaker\":\"SPEAKER_00\",\"words\":[{{\"word\":\"hello\",\"start\":0.0,\"end\":0.8,\"score\":0.9,\"speaker\":\"SPEAKER_00\"}}]}}]}}\nJSON\n",
5860 output_dir.display(),
5861 output_dir.display()
5862 ),
5863 )?;
5864 #[cfg(unix)]
5865 {
5866 use std::os::unix::fs::PermissionsExt;
5867 let mut permissions = fs::metadata(&command)?.permissions();
5868 permissions.set_mode(0o755);
5869 fs::set_permissions(&command, permissions)?;
5870 }
5871
5872 let response = transcribe(TranscriptionPipelineRequest {
5873 source: TranscriptionSource::Path {
5874 path: PathBuf::from("speech.wav"),
5875 },
5876 provider: TranscriptionProviderSelection::ExternalWhisperX(WhisperXCommandOptions {
5877 command,
5878 output_dir: Some(output_dir),
5879 ..WhisperXCommandOptions::default()
5880 }),
5881 vad: VadOptions::default(),
5882 alignment: AlignmentOptions::default(),
5883 diarization: DiarizationOptions::default(),
5884 output: TranscriptionOutputOptions::default(),
5885 })?;
5886
5887 assert!(response.accepted);
5888 assert_eq!(response.transcript.text.as_deref(), Some("hello"));
5889 assert_eq!(
5890 response.transcript.segments[0].words[0].speaker.as_deref(),
5891 Some("SPEAKER_00")
5892 );
5893 assert_eq!(response.vad_segments.len(), 1);
5894 assert_eq!(response.vad_segments[0].start_seconds, 0.29);
5895 assert_eq!(response.vad_segments[0].end_seconds, 1.47);
5896 Ok(())
5897 }
5898
5899 #[test]
5900 fn whisperx_stdout_vad_segments_parses_transcript_lines() {
5901 let segments = whisperx_stdout_vad_segments(
5902 b"noise\nTranscript: [0.29 --> 1.47] This is a test.\nTranscript: [2.00 --> 3.25] More speech.\n",
5903 );
5904
5905 assert_eq!(segments.len(), 2);
5906 assert_eq!(segments[0].start_seconds, 0.29);
5907 assert_eq!(segments[0].end_seconds, 1.47);
5908 assert_eq!(segments[1].start_seconds, 2.0);
5909 assert_eq!(segments[1].end_seconds, 3.25);
5910 }
5911
5912 #[test]
5913 fn whisperx_args_include_alignment_parity_flags() {
5914 let args = whisperx_args(
5915 Path::new("speech.wav"),
5916 Path::new("out"),
5917 &WhisperXCommandOptions {
5918 align_model: Some("facebook/wav2vec2-base-960h".to_string()),
5919 model_dir: Some(PathBuf::from("models")),
5920 model_cache_only: true,
5921 no_align: true,
5922 interpolate_method: AlignmentInterpolationMethod::Linear,
5923 return_char_alignments: true,
5924 ..WhisperXCommandOptions::default()
5925 },
5926 None,
5927 );
5928
5929 assert!(args.iter().any(|arg| arg == "--no_align"));
5930 assert!(args
5931 .windows(2)
5932 .any(|pair| pair[0] == "--align_model" && pair[1] == "facebook/wav2vec2-base-960h"));
5933 assert!(args
5934 .windows(2)
5935 .any(|pair| pair[0] == "--model_dir" && pair[1] == "models"));
5936 assert!(args.iter().any(|arg| arg == "--model_cache_only"));
5937 assert!(args
5938 .windows(2)
5939 .any(|pair| pair[0] == "--interpolate_method" && pair[1] == "linear"));
5940 assert!(args.iter().any(|arg| arg == "--return_char_alignments"));
5941 }
5942
5943 #[test]
5944 fn whisperx_args_include_task_translate() {
5945 let args = whisperx_args(
5946 Path::new("speech.wav"),
5947 Path::new("out"),
5948 &WhisperXCommandOptions {
5949 task: TranscriptionTask::Translate,
5950 ..WhisperXCommandOptions::default()
5951 },
5952 None,
5953 );
5954
5955 assert!(args
5956 .windows(2)
5957 .any(|pair| pair[0] == "--task" && pair[1] == "translate"));
5958 }
5959
5960 #[test]
5961 fn timeout_returns_typed_error() -> std::result::Result<(), Box<dyn std::error::Error>> {
5962 let temp = tempfile::tempdir()?;
5963 let command = temp.path().join("slow-whisperx.sh");
5964 fs::write(&command, "#!/usr/bin/env bash\nsleep 2\n")?;
5965 #[cfg(unix)]
5966 {
5967 use std::os::unix::fs::PermissionsExt;
5968 let mut permissions = fs::metadata(&command)?.permissions();
5969 permissions.set_mode(0o755);
5970 fs::set_permissions(&command, permissions)?;
5971 }
5972
5973 let result = transcribe(TranscriptionPipelineRequest {
5974 source: TranscriptionSource::Path {
5975 path: PathBuf::from("speech.wav"),
5976 },
5977 provider: TranscriptionProviderSelection::ExternalWhisperX(WhisperXCommandOptions {
5978 command,
5979 timeout_seconds: Some(1),
5980 ..WhisperXCommandOptions::default()
5981 }),
5982 vad: VadOptions::default(),
5983 alignment: AlignmentOptions::default(),
5984 diarization: DiarizationOptions::default(),
5985 output: TranscriptionOutputOptions::default(),
5986 });
5987 assert!(result.unwrap_err().to_string().contains("timeout"));
5988 Ok(())
5989 }
5990
5991 #[test]
5992 #[ignore]
5993 fn native_whisper_hf_cache_smoke() {
5994 if std::env::var("RUN_NATIVE_WHISPER_MODEL_CACHE_TESTS").as_deref() != Ok("1") {
5995 eprintln!(
5996 "skipping native Whisper HF cache smoke; set RUN_NATIVE_WHISPER_MODEL_CACHE_TESTS=1"
5997 );
5998 return;
5999 }
6000 #[cfg(not(all(feature = "candle", feature = "model-bundles")))]
6001 panic!("native Whisper HF cache smoke requires candle,model-bundles features");
6002
6003 #[cfg(all(feature = "candle", feature = "model-bundles"))]
6004 {
6005 let audio_path = std::env::var_os("TRANSCRIPTION_AUDIO_PATH")
6006 .map(PathBuf::from)
6007 .expect("TRANSCRIPTION_AUDIO_PATH is required");
6008 let model_dir = std::env::var_os("TRANSCRIPTION_MODEL_DIR")
6009 .map(PathBuf::from)
6010 .expect("TRANSCRIPTION_MODEL_DIR is required");
6011 let model_id =
6012 std::env::var("TRANSCRIPTION_MODEL_ID").unwrap_or_else(|_| "tiny.en".to_string());
6013 let response = transcribe(TranscriptionPipelineRequest {
6014 source: TranscriptionSource::Path { path: audio_path },
6015 provider: TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions {
6016 model_id,
6017 device: NativeDevicePreference::Cpu,
6018 language: Some("en".to_string()),
6019 model_dir: Some(model_dir),
6020 model_cache_only: true,
6021 ..CandleWhisperOptions::default()
6022 }),
6023 vad: VadOptions::default(),
6024 alignment: AlignmentOptions {
6025 enabled: false,
6026 ..AlignmentOptions::default()
6027 },
6028 diarization: DiarizationOptions::default(),
6029 output: TranscriptionOutputOptions::default(),
6030 })
6031 .expect("native Candle Whisper HF cache transcription should run");
6032 eprintln!("{}", response.diagnostics.join("\n"));
6033 assert!(response.accepted);
6034 assert!(response
6035 .diagnostics
6036 .iter()
6037 .any(|item| item == "asrModelSource=hugging-face-cache"));
6038 assert!(response
6039 .diagnostics
6040 .iter()
6041 .any(|item| item.starts_with("asrModelResolved=")));
6042 }
6043 }
6044
6045 #[test]
6046 #[ignore]
6047 fn candle_whisper_cuda_smoke_when_requested() {
6048 if std::env::var("RUN_NATIVE_TRANSCRIPTION_TESTS").as_deref() != Ok("1") {
6049 eprintln!("skipping native transcription smoke; set RUN_NATIVE_TRANSCRIPTION_TESTS=1");
6050 return;
6051 }
6052 #[cfg(not(all(feature = "candle", feature = "cuda", feature = "model-bundles")))]
6053 panic!("native transcription smoke requires candle,cuda,model-bundles features");
6054
6055 #[cfg(all(feature = "candle", feature = "cuda", feature = "model-bundles"))]
6056 {
6057 let bundle = std::env::var_os("TRANSCRIPTION_MODEL_BUNDLE")
6058 .map(PathBuf::from)
6059 .expect("TRANSCRIPTION_MODEL_BUNDLE is required");
6060 let audio_path = std::env::var_os("TRANSCRIPTION_AUDIO_PATH")
6061 .map(PathBuf::from)
6062 .expect("TRANSCRIPTION_AUDIO_PATH is required");
6063 let response = transcribe(TranscriptionPipelineRequest {
6064 source: TranscriptionSource::Path { path: audio_path },
6065 provider: TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions {
6066 model_id: "openai/whisper-tiny".to_string(),
6067 device: NativeDevicePreference::Cuda,
6068 language: Some("en".to_string()),
6069 model_bundle: Some(bundle),
6070 ..CandleWhisperOptions::default()
6071 }),
6072 vad: VadOptions::default(),
6073 alignment: AlignmentOptions::default(),
6074 diarization: DiarizationOptions::default(),
6075 output: TranscriptionOutputOptions::default(),
6076 })
6077 .expect("native Candle Whisper CUDA transcription should run");
6078 eprintln!("{}", response.diagnostics.join("\n"));
6079 assert!(response.accepted);
6080 assert!(response
6081 .diagnostics
6082 .iter()
6083 .any(|item| item == "provider=candle-whisper"));
6084 assert!(response
6085 .diagnostics
6086 .iter()
6087 .any(|item| item == "device=cuda:0"));
6088 assert!(response
6089 .diagnostics
6090 .iter()
6091 .any(|item| item.starts_with("modelId=")));
6092 assert!(response
6093 .diagnostics
6094 .iter()
6095 .any(|item| item.starts_with("bundle=")));
6096 assert!(response
6097 .diagnostics
6098 .iter()
6099 .any(|item| item.starts_with("chunkCount=")));
6100 assert!(response.diagnostics.iter().any(|item| item == "cuda=true"));
6101 assert!(
6102 response
6103 .transcript
6104 .text
6105 .as_deref()
6106 .is_some_and(|text| !text.trim().is_empty())
6107 || !response.transcript.segments.is_empty()
6108 );
6109 }
6110 }
6111
6112 #[test]
6113 #[ignore]
6114 fn candle_whisper_cuda_translate_smoke_when_requested() {
6115 if std::env::var("RUN_NATIVE_TRANSLATION_TESTS").as_deref() != Ok("1") {
6116 eprintln!("skipping native translation smoke; set RUN_NATIVE_TRANSLATION_TESTS=1");
6117 return;
6118 }
6119 #[cfg(not(all(feature = "candle", feature = "cuda", feature = "model-bundles")))]
6120 panic!("native translation smoke requires candle,cuda,model-bundles features");
6121
6122 #[cfg(all(feature = "candle", feature = "cuda", feature = "model-bundles"))]
6123 {
6124 let bundle = std::env::var_os("TRANSCRIPTION_MODEL_BUNDLE")
6125 .map(PathBuf::from)
6126 .expect("TRANSCRIPTION_MODEL_BUNDLE is required");
6127 let audio_path = std::env::var_os("TRANSCRIPTION_AUDIO_PATH")
6128 .map(PathBuf::from)
6129 .expect("TRANSCRIPTION_AUDIO_PATH is required");
6130 let response = transcribe(TranscriptionPipelineRequest {
6131 source: TranscriptionSource::Path { path: audio_path },
6132 provider: TranscriptionProviderSelection::CandleWhisper(CandleWhisperOptions {
6133 model_id: "openai/whisper-tiny".to_string(),
6134 task: TranscriptionTask::Translate,
6135 device: NativeDevicePreference::Cuda,
6136 model_bundle: Some(bundle),
6137 ..CandleWhisperOptions::default()
6138 }),
6139 vad: VadOptions::default(),
6140 alignment: AlignmentOptions {
6141 enabled: false,
6142 ..AlignmentOptions::default()
6143 },
6144 diarization: DiarizationOptions::default(),
6145 output: TranscriptionOutputOptions::default(),
6146 })
6147 .expect("native Candle Whisper CUDA translation should run");
6148 eprintln!("{}", response.diagnostics.join("\n"));
6149 assert!(response.accepted);
6150 assert!(response
6151 .diagnostics
6152 .iter()
6153 .any(|item| item == "provider=candle-whisper"));
6154 assert!(response
6155 .diagnostics
6156 .iter()
6157 .any(|item| item == "device=cuda:0"));
6158 assert!(response
6159 .diagnostics
6160 .iter()
6161 .any(|item| item == "asrTask=translate"));
6162 assert!(response
6163 .diagnostics
6164 .iter()
6165 .any(|item| item == "translationRuntime=whisper-task"));
6166 assert_eq!(response.transcript.language.as_deref(), Some("en"));
6167 assert!(
6168 response
6169 .transcript
6170 .text
6171 .as_deref()
6172 .is_some_and(|text| !text.trim().is_empty())
6173 || !response.transcript.segments.is_empty()
6174 );
6175 }
6176 }
6177
6178 #[test]
6179 #[ignore]
6180 fn ctc_alignment_wav2vec2_smoke_when_requested() {
6181 if matches!(
6182 std::env::var("RUN_NATIVE_ALIGNMENT_TESTS").as_deref(),
6183 Ok("0" | "false" | "FALSE")
6184 ) {
6185 eprintln!("skipping native alignment smoke; RUN_NATIVE_ALIGNMENT_TESTS disables it");
6186 return;
6187 }
6188 #[cfg(not(all(feature = "candle", feature = "alignment", feature = "model-bundles")))]
6189 panic!("native alignment smoke requires candle,alignment,model-bundles features");
6190
6191 #[cfg(all(feature = "candle", feature = "alignment", feature = "model-bundles"))]
6192 {
6193 let temp = tempfile::tempdir().expect("alignment smoke tempdir should be created");
6194 let (bundle, bundle_source) = default_alignment_smoke_bundle(temp.path())
6195 .expect("alignment smoke should resolve a default wav2vec2 bundle");
6196 let audio_path = default_alignment_smoke_audio_path(temp.path())
6197 .expect("alignment smoke should resolve a default WAV path");
6198 let layout = native_wav2vec2::inspect_wav2vec2_bundle_layout(&bundle)
6199 .expect("alignment smoke should inspect wav2vec2 bundle layout");
6200 eprintln!(
6201 "alignment smoke defaults: bundleSource={bundle_source} bundle={} audio={}",
6202 bundle.display(),
6203 audio_path.display()
6204 );
6205 eprintln!("wav2vec2 layout report: {layout:#?}");
6206 let transcript_text = std::env::var("ALIGNMENT_TRANSCRIPT_TEXT")
6207 .unwrap_or_else(|_| "hello world".to_string());
6208 let audio = LoadedAudio::mono_16khz_from_source(&TranscriptionSource::Path {
6209 path: audio_path,
6210 })
6211 .expect("alignment smoke requires readable WAV audio");
6212 let mut segment = TranscriptSegmentContract::new(0, transcript_text.clone());
6213 segment.start_seconds = Some(0.0);
6214 segment.end_seconds = Some(audio.duration_seconds().clamp(1.0 / 16_000.0, 1.0));
6215 let transcript =
6216 TranscriptionContract::from_segments(None, Some("en".to_string()), vec![segment])
6217 .expect("alignment smoke transcript should validate");
6218 let request = AlignmentRequest {
6219 audio,
6220 transcript,
6221 language: Some("en".to_string()),
6222 model_id: default_alignment_model(),
6223 };
6224 let mut aligner = CtcForcedAligner {
6225 options: AlignmentOptions {
6226 enabled: true,
6227 model_bundle: Some(bundle),
6228 return_char_alignments: true,
6229 ..AlignmentOptions::default()
6230 },
6231 };
6232 let response = aligner
6233 .align(request)
6234 .expect("native wav2vec2/CTC alignment should run");
6235 eprintln!("{}", response.diagnostics.join("\n"));
6236 assert!(!response.words.is_empty());
6237 assert!(response
6238 .words
6239 .iter()
6240 .all(|word| word.end_seconds >= word.start_seconds));
6241 assert!(response.words.iter().all(|word| word.confidence.is_some()));
6242 assert!(!response.chars.is_empty());
6243 assert!(response
6244 .diagnostics
6245 .iter()
6246 .any(|item| item == "alignmentProvider=ctc-forced-aligner"));
6247 assert!(response
6248 .diagnostics
6249 .iter()
6250 .any(|item| item == "alignmentModelExecution=candle-wav2vec2"));
6251 assert!(response
6252 .diagnostics
6253 .iter()
6254 .any(|item| item == "alignmentModelSource=explicit-bundle"));
6255 assert!(response
6256 .diagnostics
6257 .iter()
6258 .any(|item| item == "alignmentInterpolateMethod=nearest"));
6259 assert!(response
6260 .diagnostics
6261 .iter()
6262 .any(|item| item == "returnCharAlignments=true"));
6263 }
6264 }
6265}