Skip to main content

audio_analysis_transcription/
lib.rs

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
46/// Backward-compatible name for Transcript Speaker Assignment policy.
47pub type SpeakerAssignmentPolicy = SpeakerTranscriptAssignmentPolicy;
48
49/// Request for an audio/video-to-text transcription pipeline.
50#[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/// Phase-level observer event emitted by the native transcription pipeline.
66#[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
118/// Observer for phase-level native transcription progress.
119pub 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    /// Returns whether the current pipeline should stop at its next safe phase boundary.
145    fn cancellation_requested(&self) -> bool {
146        false
147    }
148}
149
150/// Observer implementation that discards all events.
151#[derive(Debug, Default)]
152pub struct NoopTranscriptionPipelineObserver;
153
154impl TranscriptionPipelineObserver for NoopTranscriptionPipelineObserver {
155    fn observe(&mut self, _event: TranscriptionPipelineEvent) {}
156}
157
158/// Source accepted by transcription providers.
159#[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/// Provider selection.
187#[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/// Whisper speech task.
225#[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/// Options for the Candle Whisper native provider.
250#[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/// Native Candle Whisper compute-type preference.
300#[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/// Native Candle Whisper chunk decode runtime.
340#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
341#[serde(rename_all = "camelCase")]
342pub enum CandleWhisperDecodeRuntime {
343    /// Existing safe per-window autoregressive decode with KV-cache reuse inside each window.
344    #[default]
345    AutoregressiveKvCache,
346    /// Future true tensor-batched active-row decode path.
347    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/// Options for native whisper.cpp compatibility.
367#[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/// Native execution device preference.
396#[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/// Options for the external WhisperX command provider.
406#[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
448/// Backward-compatible name for the external WhisperX provider options.
449pub 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/// WhisperX device selection.
479#[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/// Timestamp interpolation behavior for missing alignment spans.
496#[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/// VAD options used before ASR chunking.
516#[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/// Forced-alignment options.
578#[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/// Diarization options.
623#[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/// Output preferences for transcription.
647#[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/// Artifact produced or discovered by a provider.
667#[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/// Response from a transcription pipeline.
675#[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/// Metadata-only provider plan.
691#[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/// Loaded audio for native providers.
703#[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/// A speech activity span used for ASR chunking.
727#[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/// ASR provider request.
765#[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/// ASR provider response.
777#[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/// Forced-alignment request.
788#[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/// Forced-alignment response.
798#[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/// One aligned word timing.
810#[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/// One aligned character timing.
823#[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/// Alignment summary included in pipeline responses.
838#[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/// VAD request.
847#[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/// VAD response.
855#[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
863/// Trait for audio transcription providers.
864pub 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
877/// Trait for forced alignment providers.
878pub 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
891/// Trait for transcription VAD providers.
892pub trait TranscriptionVadProvider {
893    fn provider_id(&self) -> &str;
894    fn detect_speech(&mut self, request: VadRequest) -> Result<VadResponse>;
895}
896
897/// Trait for transcript diarization and speaker assignment providers.
898pub 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/// Native deterministic speaker diarization adapter.
909#[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/// Native pyannote community diarization adapter.
987#[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/// Default pure-Rust energy VAD provider.
1073#[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/// Feature-gated Candle Whisper provider.
1091#[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/// Candle Whisper provider that keeps a compatible native model session loaded
1186/// across transcription requests.
1187///
1188/// Downstream callers can use this as the public provider-reuse primitive with
1189/// `run_transcription_pipeline_with_observer`. Compatible repeated requests
1190/// emit `TranscriptionPipelineEvent::ModelReuse` and include response
1191/// diagnostics such as `asrModelSession=reused` when a loaded session is reused.
1192#[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/// Native whisper.cpp compatibility provider.
1334#[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/// Feature-gated CTC forced aligner.
1361#[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/// External command provider for Python WhisperX.
1398#[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/// Native runner configuration for repeated finite transcription requests.
1419///
1420/// The options describe the provider stack the runner owns. Requests passed to
1421/// [`NativeTranscriptionRunner::run`] must use the same provider, VAD,
1422/// alignment, and diarization options so the runner can safely reuse provider
1423/// state across different sources without changing pipeline semantics.
1424#[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    /// Builds runner options from the reusable parts of a pipeline request.
1438    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
1459/// Owning native transcription runner for repeated finite requests.
1460///
1461/// Use this when callers want the crate to own native provider/session
1462/// lifecycle across multiple compatible requests. Advanced callers can still
1463/// call [`run_transcription_pipeline_with_observer`] directly with their own
1464/// provider trait implementations.
1465pub 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    /// Builds the default native provider stack for the configured provider.
1482    ///
1483    /// Candle Whisper requests use [`ReusableCandleWhisperTranscriber`] so
1484    /// compatible repeated requests can report reuse through
1485    /// [`TranscriptionPipelineEvent::ModelReuse`] and response diagnostics such
1486    /// as `asrModelSession=reused`.
1487    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    /// Builds a runner from caller-provided provider adapters.
1537    ///
1538    /// This keeps the customization seam at the existing provider traits rather
1539    /// than introducing a separate test-only abstraction.
1540    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    /// Runs a request through the owned provider stack.
1558    ///
1559    /// Default runners rebuild their crate-owned providers when request
1560    /// options change. Runners built from caller-provided adapters still
1561    /// require exact option compatibility because the runner cannot rebuild
1562    /// external provider state.
1563    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
1649/// Runs a transcription request with the selected primary provider.
1650pub 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
1695/// Runs the provider-agnostic native transcription pipeline.
1696pub 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
1714/// Runs the provider-agnostic native transcription pipeline and emits phase events.
1715pub 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
1901/// Parses existing WhisperX JSON without running external tools.
1902pub 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
1907/// Returns provider plans.
1908pub 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
1916/// Returns the primary Candle Whisper provider plan.
1917pub 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
1937/// Returns the whisper.cpp compatibility provider plan.
1938pub 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
1953/// Returns the external WhisperX provider plan.
1954pub 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}