use std::collections::HashMap;
use std::sync::Arc;
use crate::inference::{EndpointMode, EndpointReason, Engine, PRED_HIDDEN, WordInfo};
use crate::runtime::mock::{MockFactory, MockSession};
use crate::runtime::tensor::{Shape, Tensor, TensorData};
const ENC_DIM: usize = 768;
fn tiny_mock_engine() -> (Engine, tempfile::TempDir) {
let tmp = tempfile::tempdir().expect("tempdir");
let dir = tmp.path();
std::fs::write(dir.join("v3_rnnt_encoder_int8.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_rnnt_decoder.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_rnnt_joint.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_vocab.txt"), "\u{2581}hi\n<blk>\n").unwrap();
let mut sessions: HashMap<String, Arc<MockSession>> = HashMap::new();
sessions.insert(
"v3_rnnt_encoder_int8".into(),
Arc::new(MockSession::new(
vec![Shape::new(vec![1, 64, 1]), Shape::new(vec![1])],
vec![
Tensor::new(
Shape::new(vec![1, ENC_DIM, 1]),
TensorData::F32(vec![0.0; ENC_DIM]),
)
.unwrap(),
Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![1])).unwrap(),
],
)),
);
sessions.insert(
"v3_rnnt_decoder".into(),
Arc::new(MockSession::new(
vec![
Shape::new(vec![1, 1]),
Shape::new(vec![1, 1, PRED_HIDDEN]),
Shape::new(vec![1, 1, PRED_HIDDEN]),
],
vec![
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
],
)),
);
sessions.insert(
"v3_rnnt_joint".into(),
Arc::new(MockSession::new(
vec![
Shape::new(vec![1, ENC_DIM, 1]),
Shape::new(vec![1, PRED_HIDDEN, 1]),
],
vec![Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![0.0; 2])).unwrap()],
)),
);
let factory = Box::new(MockFactory::new(sessions));
let engine = Engine::load_with_factory(dir, None, 1, 1, 0, factory, 1)
.expect("engine should load with mock runtime");
(engine, tmp)
}
#[test]
fn test_engine_loads_with_mock_runtime() {
let _ = tiny_mock_engine();
}
#[test]
fn test_engine_mock_runtime_decodes_silence() {
let (engine, _tmp) = tiny_mock_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let samples = vec![0.0f32; 100]; let result = engine
.transcribe_samples(&samples, &mut guard)
.expect("mock decode must not error");
assert!(result.text.is_empty(), "blank-only decode yields no text");
assert!(result.words.is_empty());
assert!((result.duration_s - 100.0 / 16000.0).abs() < 1e-9);
}
#[test]
fn test_validate_overrides_truth_table() {
use crate::inference::{
HotwordOverride, MAX_HOTWORD_PHRASE_CHARS, MAX_HOTWORDS_PER_REQUEST, OverrideError,
TranscribeOverrides,
};
let (engine, _tmp) = tiny_mock_engine();
assert!(!engine.has_vad(), "mock engine has no VAD");
assert!(!engine.has_punctuator(), "mock engine has no punctuator");
assert_eq!(
engine.validate_overrides(&TranscribeOverrides::default()),
Ok(())
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides {
vad: Some(true),
..Default::default()
}),
Err(OverrideError::VadNotLoaded)
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides {
vad: Some(false),
..Default::default()
}),
Ok(())
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides {
punctuation: Some(true),
..Default::default()
}),
Err(OverrideError::PunctuationNotAvailable)
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides {
punctuation: Some(false),
..Default::default()
}),
Ok(())
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides {
itn: Some(true),
..Default::default()
}),
Ok(())
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides {
itn: Some(false),
..Default::default()
}),
Ok(())
);
use crate::inference::HotwordError;
assert_eq!(
engine.validate_hotwords(&HotwordOverride::new(vec![], None)),
Ok(())
);
assert_eq!(
engine.validate_hotwords(&HotwordOverride::new(
vec!["ok".into(), "fine".into()],
Some(3.0),
)),
Ok(())
);
let at_cap: Vec<String> = (0..MAX_HOTWORDS_PER_REQUEST)
.map(|i| format!("w{i}"))
.collect();
assert_eq!(
engine.validate_hotwords(&HotwordOverride::new(at_cap, None)),
Ok(())
);
let over_cap: Vec<String> = (0..=MAX_HOTWORDS_PER_REQUEST)
.map(|i| format!("w{i}"))
.collect();
assert_eq!(
engine.validate_hotwords(&HotwordOverride::new(over_cap, None)),
Err(HotwordError::TooManyHotwords)
);
let ok_phrase: String = "а".repeat(MAX_HOTWORD_PHRASE_CHARS);
assert_eq!(
engine.validate_hotwords(&HotwordOverride::new(vec![ok_phrase], None)),
Ok(())
);
let long_phrase: String = "а".repeat(MAX_HOTWORD_PHRASE_CHARS + 1);
assert_eq!(
engine.validate_hotwords(&HotwordOverride::new(vec![long_phrase], None)),
Err(HotwordError::PhraseTooLong)
);
}
#[test]
fn test_request_hotword_biaser_semantics() {
use crate::inference::{HotwordOverride, TranscribeOverrides};
let (engine, _tmp) = tiny_mock_engine();
let engine = engine.with_hotwords(&[("hi".into(), 1.0)], 5.0);
assert!(engine.has_hotwords(), "boot biaser attached");
let off = HotwordOverride::new(vec![], None);
assert!(
engine.build_request_biaser(&off).is_none(),
"empty override forces biasing off"
);
let on = HotwordOverride::new(vec!["hi".into()], Some(7.0));
let built = engine
.build_request_biaser(&on)
.expect("representable phrase should compile");
assert_eq!(built.phrase_count(), 1);
let default_boost = HotwordOverride::new(vec!["hi".into()], None);
assert!(engine.build_request_biaser(&default_boost).is_some());
let junk = HotwordOverride::new(vec!["яяяя".into()], None);
assert!(
engine.build_request_biaser(&junk).is_none(),
"unrepresentable phrases drop the temporary biaser"
);
assert_eq!(
engine.validate_overrides(&TranscribeOverrides::default()),
Ok(())
);
assert!(engine.has_hotwords());
}
#[test]
fn test_ctc_head_attaches_a_biaser() {
use crate::inference::HotwordOverride;
use crate::model::ModelVariant;
let (mut engine, _tmp) = tiny_mock_engine();
engine.variant = ModelVariant::MlCtc;
let engine = engine.with_hotwords(&[("hi".into(), 1.0)], 5.0);
assert!(
engine.has_hotwords(),
"a CTC head biases through the prefix beam"
);
let per_request = HotwordOverride::new(vec!["hi".into()], Some(7.0));
assert!(
engine.build_request_biaser(&per_request).is_some(),
"a per-request glossary applies on a CTC head too"
);
}
#[test]
fn test_transcribe_samples_with_overrides_vad_off_matches_default() {
use crate::inference::TranscribeOverrides;
let (engine, _tmp) = tiny_mock_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let samples = vec![0.0f32; 100];
let baseline = engine
.transcribe_samples(&samples, &mut guard)
.expect("baseline decode");
let with_vad_off = engine
.transcribe_samples_with_overrides(
&samples,
&mut guard,
&TranscribeOverrides {
vad: Some(false),
..Default::default()
},
None,
false,
None,
super::super::DecodeControls::default(),
)
.expect("vad-off decode");
assert_eq!(baseline.text, with_vad_off.text);
assert_eq!(baseline.words.len(), with_vad_off.words.len());
}
#[test]
fn test_diarization_requested_without_speaker_model_records_notice() {
use crate::inference::{DiarizationOutcome, TranscribeRequest, TranscribeSource};
use std::sync::{Arc, OnceLock};
let (engine, _tmp) = tiny_mock_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let samples = vec![0.0f32; 100];
let sink: Arc<OnceLock<DiarizationOutcome>> = Arc::new(OnceLock::new());
let req = TranscribeRequest::new(TranscribeSource::Samples(&samples))
.with_diarization(true)
.with_diarization_outcome(Some(sink.clone()));
engine
.transcribe_request(req, &mut guard)
.expect("decode should succeed");
assert_eq!(
sink.get().copied(),
Some(DiarizationOutcome::NoSpeakerModel),
"diarization requested without a model must be reported, not silent"
);
}
#[test]
fn test_create_state_postprocess_overrides_default_to_none() {
let (engine, _tmp) = tiny_mock_engine();
let state = engine.create_state(false);
assert_eq!(state.punctuation, None);
assert_eq!(state.itn, None);
}
#[test]
fn test_flush_state_applies_itn_per_engine_default() {
let (engine, _tmp) = tiny_mock_engine();
let engine = engine.with_itn(true);
let mut state = engine.create_state(false);
state.assembler.set_words(vec![
super::word("двадцать", 0.0, 0.4),
super::word("один", 0.4, 0.8),
]);
let seg = engine.flush_state(&mut state).expect("flush");
assert_eq!(seg.text, "21");
assert_eq!(seg.words[0].word, "двадцать");
assert_eq!(seg.words[1].word, "один");
}
#[test]
fn test_flush_state_session_override_disables_itn() {
let (engine, _tmp) = tiny_mock_engine();
let engine = engine.with_itn(true);
let mut state = engine.create_state(false);
state.itn = Some(false);
state.assembler.set_words(vec![
super::word("двадцать", 0.0, 0.4),
super::word("один", 0.4, 0.8),
]);
let seg = engine.flush_state(&mut state).expect("flush");
assert_eq!(seg.text, "двадцать один");
}
#[test]
fn test_flush_state_default_leaves_text_raw() {
let (engine, _tmp) = tiny_mock_engine();
let mut state = engine.create_state(false);
state.assembler.set_words(vec![
super::word("двадцать", 0.0, 0.4),
super::word("один", 0.4, 0.8),
]);
let seg = engine.flush_state(&mut state).expect("flush");
assert_eq!(seg.text, "двадцать один");
}
#[test]
fn test_flush_state_punctuation_request_without_punctuator_is_noop() {
let (engine, _tmp) = tiny_mock_engine();
let mut state = engine.create_state(false);
state.punctuation = Some(true);
state
.assembler
.set_words(vec![super::word("hello", 0.0, 0.4)]);
let seg = engine.flush_state(&mut state).expect("flush");
assert_eq!(seg.text, "hello");
}
fn blank_run_engine() -> (Engine, tempfile::TempDir) {
let tmp = tempfile::tempdir().expect("tempdir");
let dir = tmp.path();
std::fs::write(dir.join("v3_rnnt_encoder_int8.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_rnnt_decoder.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_rnnt_joint.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_vocab.txt"), "\u{2581}hi\n<blk>\n").unwrap();
const MEL_FRAMES: usize = 79;
const ENC_LEN: usize = 16;
let mut sessions: HashMap<String, Arc<MockSession>> = HashMap::new();
sessions.insert(
"v3_rnnt_encoder_int8".into(),
Arc::new(MockSession::new(
vec![Shape::new(vec![1, 64, MEL_FRAMES]), Shape::new(vec![1])],
vec![
Tensor::new(
Shape::new(vec![1, ENC_DIM, ENC_LEN]),
TensorData::F32(vec![0.0; ENC_DIM * ENC_LEN]),
)
.unwrap(),
Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![ENC_LEN as i64])).unwrap(),
],
)),
);
sessions.insert(
"v3_rnnt_decoder".into(),
Arc::new(MockSession::new(
vec![
Shape::new(vec![1, 1]),
Shape::new(vec![1, 1, PRED_HIDDEN]),
Shape::new(vec![1, 1, PRED_HIDDEN]),
],
vec![
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
],
)),
);
sessions.insert(
"v3_rnnt_joint".into(),
Arc::new(
MockSession::new(
vec![
Shape::new(vec![1, ENC_DIM, 1]),
Shape::new(vec![1, PRED_HIDDEN, 1]),
],
vec![
Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![0.0; 2])).unwrap(),
],
)
.with_script(vec![
vec![
Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![2.0, 0.0]))
.unwrap(),
],
vec![
Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![0.0, 2.0]))
.unwrap(),
],
]),
),
);
let factory = Box::new(MockFactory::new(sessions));
let engine = Engine::load_with_factory(dir, None, 1, 1, 0, factory, 1)
.expect("engine should load with mock runtime");
(engine, tmp)
}
#[test]
fn test_process_chunk_blank_endpoint_finalizes_without_vad() {
let (engine, _tmp) = blank_run_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let mut state = engine.create_state(false);
let chunk = vec![0.0f32; 12800]; let segs = engine
.process_chunk(&chunk, &mut state, &mut guard)
.expect("mock decode must not error");
assert_eq!(
segs.len(),
1,
"blank-run endpoint must finalize the segment"
);
assert!(segs[0].is_final);
assert!(segs[0].speech_final);
assert_eq!(segs[0].endpoint_reason, Some(EndpointReason::Blank));
assert_eq!(segs[0].text, "hi");
}
#[test]
fn test_process_chunk_blank_endpoint_ignored_with_vad() {
let (engine, _tmp) = blank_run_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let mut state = engine.create_state(false);
state.vad_endpointer = Some(crate::vad::VadEndpointer::new(
&crate::vad::VadConfig::default(),
));
let chunk = vec![0.0f32; 12800];
let segs = engine
.process_chunk(&chunk, &mut state, &mut guard)
.expect("mock decode must not error");
assert_eq!(segs.len(), 1, "decoded words still surface as a partial");
assert!(
!segs[0].is_final,
"blank-run must not finalize while a VAD owns endpointing"
);
assert!(!segs[0].speech_final);
assert!(segs[0].endpoint_reason.is_none());
assert_eq!(segs[0].text, "hi");
}
fn blank_run_engine_window_cap() -> (Engine, tempfile::TempDir) {
let tmp = tempfile::tempdir().expect("tempdir");
let dir = tmp.path();
std::fs::write(dir.join("v3_rnnt_encoder_int8.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_rnnt_decoder.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_rnnt_joint.onnx"), b"").unwrap();
std::fs::write(dir.join("v3_vocab.txt"), "\u{2581}hi\n<blk>\n").unwrap();
const MEL_FRAMES: usize = 249;
const ENC_LEN: usize = 16;
let mut sessions: HashMap<String, Arc<MockSession>> = HashMap::new();
sessions.insert(
"v3_rnnt_encoder_int8".into(),
Arc::new(MockSession::new(
vec![Shape::new(vec![1, 64, MEL_FRAMES]), Shape::new(vec![1])],
vec![
Tensor::new(
Shape::new(vec![1, ENC_DIM, ENC_LEN]),
TensorData::F32(vec![0.0; ENC_DIM * ENC_LEN]),
)
.unwrap(),
Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![ENC_LEN as i64])).unwrap(),
],
)),
);
sessions.insert(
"v3_rnnt_decoder".into(),
Arc::new(MockSession::new(
vec![
Shape::new(vec![1, 1]),
Shape::new(vec![1, 1, PRED_HIDDEN]),
Shape::new(vec![1, 1, PRED_HIDDEN]),
],
vec![
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
Tensor::new(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
)
.unwrap(),
],
)),
);
sessions.insert(
"v3_rnnt_joint".into(),
Arc::new(
MockSession::new(
vec![
Shape::new(vec![1, ENC_DIM, 1]),
Shape::new(vec![1, PRED_HIDDEN, 1]),
],
vec![
Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![0.0; 2])).unwrap(),
],
)
.with_script(vec![
vec![
Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![2.0, 0.0]))
.unwrap(),
],
vec![
Tensor::new(Shape::new(vec![1, 1, 2]), TensorData::F32(vec![0.0, 2.0]))
.unwrap(),
],
]),
),
);
let factory = Box::new(MockFactory::new(sessions));
let engine = Engine::load_with_factory(dir, None, 1, 1, 0, factory, 1)
.expect("engine should load with mock runtime");
(engine, tmp)
}
#[test]
fn test_process_chunk_window_cap_emits_partial_not_final() {
let (engine, _tmp) = blank_run_engine_window_cap();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let mut state = engine.create_state(false);
state.vad_endpointer = Some(crate::vad::VadEndpointer::new(
&crate::vad::VadConfig::default(),
));
let chunk = vec![0.0f32; 16000 * 5 / 2];
let segs = engine
.process_chunk(&chunk, &mut state, &mut guard)
.expect("mock decode must not error");
assert_eq!(segs.len(), 1, "cap must still surface decoded text");
assert!(
!segs[0].is_final,
"window cap must not emit is_final (got final text={:?})",
segs[0].text
);
assert!(!segs[0].speech_final);
assert!(segs[0].endpoint_reason.is_none());
assert_eq!(segs[0].text, "hi");
assert!(
!state.assembler.is_empty(),
"stable prefix must remain after cap commit"
);
}
#[test]
fn test_process_chunk_assistant_mode_ignores_blank_without_vad() {
let (engine, _tmp) = blank_run_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let mut state = engine.create_state(false);
state.endpoint_mode = EndpointMode::Assistant;
let chunk = vec![0.0f32; 12800];
let segs = engine
.process_chunk(&chunk, &mut state, &mut guard)
.expect("mock decode must not error");
assert_eq!(segs.len(), 1);
assert!(
!segs[0].is_final,
"assistant mode must not finalize on blank-run alone"
);
assert_eq!(segs[0].text, "hi");
}
#[test]
fn test_process_chunk_manual_mode_never_auto_finalizes() {
let (engine, _tmp) = blank_run_engine();
let mut guard = engine.pool.checkout_blocking().expect("checkout");
let mut state = engine.create_state(false);
state.endpoint_mode = EndpointMode::Manual;
let chunk = vec![0.0f32; 12800];
let segs = engine
.process_chunk(&chunk, &mut state, &mut guard)
.expect("mock decode must not error");
assert_eq!(segs.len(), 1);
assert!(!segs[0].is_final);
}
#[test]
fn test_flush_state_marks_stop_endpoint_reason() {
let (engine, _tmp) = blank_run_engine();
let mut state = engine.create_state(false);
state
.assembler
.append(vec![WordInfo::new("bye", 0.0, 0.3, 0.9, None)]);
let seg = engine.flush_state(&mut state).expect("flush");
assert!(seg.is_final);
assert!(seg.speech_final);
assert_eq!(seg.endpoint_reason, Some(EndpointReason::Stop));
assert_eq!(seg.text, "bye");
}