use wasm_bindgen::prelude::*;
#[wasm_bindgen]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WorkerMessageType {
Init,
Transcribe,
DetectLanguage,
Abort,
Ready,
Result,
Language,
Progress,
Error,
}
impl WorkerMessageType {
#[must_use]
#[allow(clippy::should_implement_trait)]
pub fn from_str(s: &str) -> Option<Self> {
match s {
"init" => Some(Self::Init),
"transcribe" => Some(Self::Transcribe),
"detectLanguage" => Some(Self::DetectLanguage),
"abort" => Some(Self::Abort),
"ready" => Some(Self::Ready),
"result" => Some(Self::Result),
"language" => Some(Self::Language),
"progress" => Some(Self::Progress),
"error" => Some(Self::Error),
_ => None,
}
}
#[must_use]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Init => "init",
Self::Transcribe => "transcribe",
Self::DetectLanguage => "detectLanguage",
Self::Abort => "abort",
Self::Ready => "ready",
Self::Result => "result",
Self::Language => "language",
Self::Progress => "progress",
Self::Error => "error",
}
}
}
#[wasm_bindgen]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum WorkerState {
#[default]
Uninitialized,
Loading,
Ready,
Transcribing,
DetectingLanguage,
Aborted,
Error,
}
impl WorkerState {
#[must_use]
pub fn is_ready(&self) -> bool {
matches!(self, Self::Ready)
}
#[must_use]
pub fn is_busy(&self) -> bool {
matches!(
self,
Self::Transcribing | Self::DetectingLanguage | Self::Loading
)
}
#[must_use]
pub const fn name(&self) -> &'static str {
match self {
Self::Uninitialized => "uninitialized",
Self::Loading => "loading",
Self::Ready => "ready",
Self::Transcribing => "transcribing",
Self::DetectingLanguage => "detectingLanguage",
Self::Aborted => "aborted",
Self::Error => "error",
}
}
}
#[wasm_bindgen(js_name = workerStateIsReady)]
pub fn worker_state_is_ready(state: WorkerState) -> bool {
state.is_ready()
}
#[wasm_bindgen(js_name = workerStateIsBusy)]
pub fn worker_state_is_busy(state: WorkerState) -> bool {
state.is_busy()
}
#[wasm_bindgen(js_name = workerStateName)]
pub fn worker_state_name(state: WorkerState) -> String {
state.name().to_string()
}
#[wasm_bindgen]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProgressPhase {
LoadingModel,
Resampling,
MelSpectrogram,
Encoding,
Decoding,
Postprocessing,
Complete,
}
impl ProgressPhase {
#[must_use]
pub const fn name(&self) -> &'static str {
match self {
Self::LoadingModel => "Loading model",
Self::Resampling => "Resampling audio",
Self::MelSpectrogram => "Computing mel spectrogram",
Self::Encoding => "Running encoder",
Self::Decoding => "Running decoder",
Self::Postprocessing => "Processing results",
Self::Complete => "Complete",
}
}
#[must_use]
pub const fn weight(&self) -> u32 {
match self {
Self::LoadingModel | Self::MelSpectrogram => 10,
Self::Resampling | Self::Postprocessing => 5,
Self::Encoding => 30,
Self::Decoding => 40,
Self::Complete => 0,
}
}
}
#[wasm_bindgen]
#[derive(Debug, Clone)]
pub struct WorkerProgress {
phase: ProgressPhase,
phase_progress: f32,
overall_progress: f32,
message: String,
}
#[wasm_bindgen]
impl WorkerProgress {
#[wasm_bindgen(constructor)]
pub fn new() -> Self {
Self {
phase: ProgressPhase::LoadingModel,
phase_progress: 0.0,
overall_progress: 0.0,
message: String::new(),
}
}
#[wasm_bindgen(getter)]
pub fn phase(&self) -> ProgressPhase {
self.phase
}
#[wasm_bindgen(js_name = phaseName)]
pub fn phase_name(&self) -> String {
self.phase.name().to_string()
}
#[wasm_bindgen(getter, js_name = phaseProgress)]
pub fn phase_progress(&self) -> f32 {
self.phase_progress
}
#[wasm_bindgen(getter, js_name = overallProgress)]
pub fn overall_progress(&self) -> f32 {
self.overall_progress
}
#[wasm_bindgen(getter)]
pub fn message(&self) -> String {
self.message.clone()
}
pub fn update(&mut self, phase: ProgressPhase, phase_progress: f32, message: &str) {
self.phase = phase;
self.phase_progress = phase_progress.clamp(0.0, 100.0);
self.message = message.to_string();
self.calculate_overall();
}
pub fn next_phase(&mut self, phase: ProgressPhase) {
self.phase = phase;
self.phase_progress = 0.0;
self.message = phase.name().to_string();
self.calculate_overall();
}
fn calculate_overall(&mut self) {
let phases = [
ProgressPhase::LoadingModel,
ProgressPhase::Resampling,
ProgressPhase::MelSpectrogram,
ProgressPhase::Encoding,
ProgressPhase::Decoding,
ProgressPhase::Postprocessing,
ProgressPhase::Complete,
];
let total_weight: u32 = phases.iter().map(|p| p.weight()).sum();
let mut completed_weight: u32 = 0;
for p in &phases {
if *p == self.phase {
let phase_contribution = p.weight() as f32 * (self.phase_progress / 100.0);
completed_weight += phase_contribution as u32;
break;
}
completed_weight += p.weight();
}
self.overall_progress = (completed_weight as f32 / total_weight as f32) * 100.0;
}
pub fn complete(&mut self) {
self.phase = ProgressPhase::Complete;
self.phase_progress = 100.0;
self.overall_progress = 100.0;
self.message = "Complete".to_string();
}
}
impl Default for WorkerProgress {
fn default() -> Self {
Self::new()
}
}
#[wasm_bindgen]
#[derive(Debug, Clone)]
pub struct WorkerConfig {
model_type: String,
enable_progress: bool,
progress_interval_ms: u32,
}
#[wasm_bindgen]
impl WorkerConfig {
#[wasm_bindgen(constructor)]
pub fn new() -> Self {
Self {
model_type: "tiny".to_string(),
enable_progress: true,
progress_interval_ms: 100,
}
}
#[wasm_bindgen(setter, js_name = modelType)]
pub fn set_model_type(&mut self, model_type: &str) {
self.model_type = model_type.to_string();
}
#[wasm_bindgen(getter, js_name = modelType)]
pub fn model_type(&self) -> String {
self.model_type.clone()
}
#[wasm_bindgen(setter, js_name = enableProgress)]
pub fn set_enable_progress(&mut self, enable: bool) {
self.enable_progress = enable;
}
#[wasm_bindgen(getter, js_name = enableProgress)]
pub fn enable_progress(&self) -> bool {
self.enable_progress
}
#[wasm_bindgen(setter, js_name = progressIntervalMs)]
pub fn set_progress_interval_ms(&mut self, interval: u32) {
self.progress_interval_ms = interval;
}
#[wasm_bindgen(getter, js_name = progressIntervalMs)]
pub fn progress_interval_ms(&self) -> u32 {
self.progress_interval_ms
}
}
impl Default for WorkerConfig {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_worker_message_type_from_str() {
assert_eq!(
WorkerMessageType::from_str("init"),
Some(WorkerMessageType::Init)
);
assert_eq!(
WorkerMessageType::from_str("transcribe"),
Some(WorkerMessageType::Transcribe)
);
assert_eq!(
WorkerMessageType::from_str("detectLanguage"),
Some(WorkerMessageType::DetectLanguage)
);
assert_eq!(
WorkerMessageType::from_str("abort"),
Some(WorkerMessageType::Abort)
);
assert_eq!(
WorkerMessageType::from_str("ready"),
Some(WorkerMessageType::Ready)
);
assert_eq!(
WorkerMessageType::from_str("result"),
Some(WorkerMessageType::Result)
);
assert_eq!(
WorkerMessageType::from_str("language"),
Some(WorkerMessageType::Language)
);
assert_eq!(
WorkerMessageType::from_str("progress"),
Some(WorkerMessageType::Progress)
);
assert_eq!(
WorkerMessageType::from_str("error"),
Some(WorkerMessageType::Error)
);
assert_eq!(WorkerMessageType::from_str("invalid"), None);
}
#[test]
fn test_worker_message_type_as_str() {
assert_eq!(WorkerMessageType::Init.as_str(), "init");
assert_eq!(WorkerMessageType::Transcribe.as_str(), "transcribe");
assert_eq!(WorkerMessageType::DetectLanguage.as_str(), "detectLanguage");
assert_eq!(WorkerMessageType::Abort.as_str(), "abort");
assert_eq!(WorkerMessageType::Ready.as_str(), "ready");
assert_eq!(WorkerMessageType::Result.as_str(), "result");
assert_eq!(WorkerMessageType::Language.as_str(), "language");
assert_eq!(WorkerMessageType::Progress.as_str(), "progress");
assert_eq!(WorkerMessageType::Error.as_str(), "error");
}
#[test]
fn test_worker_message_type_roundtrip() {
let types = [
WorkerMessageType::Init,
WorkerMessageType::Transcribe,
WorkerMessageType::DetectLanguage,
WorkerMessageType::Abort,
WorkerMessageType::Ready,
WorkerMessageType::Result,
WorkerMessageType::Language,
WorkerMessageType::Progress,
WorkerMessageType::Error,
];
for t in types {
let s = t.as_str();
let parsed = WorkerMessageType::from_str(s);
assert_eq!(parsed, Some(t));
}
}
#[test]
fn test_worker_state_default() {
let state = WorkerState::default();
assert_eq!(state, WorkerState::Uninitialized);
}
#[test]
fn test_worker_state_is_ready() {
assert!(!WorkerState::Uninitialized.is_ready());
assert!(!WorkerState::Loading.is_ready());
assert!(WorkerState::Ready.is_ready());
assert!(!WorkerState::Transcribing.is_ready());
assert!(!WorkerState::Error.is_ready());
}
#[test]
fn test_worker_state_is_busy() {
assert!(!WorkerState::Uninitialized.is_busy());
assert!(WorkerState::Loading.is_busy());
assert!(!WorkerState::Ready.is_busy());
assert!(WorkerState::Transcribing.is_busy());
assert!(WorkerState::DetectingLanguage.is_busy());
assert!(!WorkerState::Error.is_busy());
}
#[test]
fn test_worker_state_name() {
assert_eq!(WorkerState::Uninitialized.name(), "uninitialized");
assert_eq!(WorkerState::Loading.name(), "loading");
assert_eq!(WorkerState::Ready.name(), "ready");
assert_eq!(WorkerState::Transcribing.name(), "transcribing");
assert_eq!(WorkerState::DetectingLanguage.name(), "detectingLanguage");
assert_eq!(WorkerState::Aborted.name(), "aborted");
assert_eq!(WorkerState::Error.name(), "error");
}
#[test]
fn test_progress_phase_name() {
assert_eq!(ProgressPhase::LoadingModel.name(), "Loading model");
assert_eq!(ProgressPhase::Resampling.name(), "Resampling audio");
assert_eq!(
ProgressPhase::MelSpectrogram.name(),
"Computing mel spectrogram"
);
assert_eq!(ProgressPhase::Encoding.name(), "Running encoder");
assert_eq!(ProgressPhase::Decoding.name(), "Running decoder");
assert_eq!(ProgressPhase::Postprocessing.name(), "Processing results");
assert_eq!(ProgressPhase::Complete.name(), "Complete");
}
#[test]
fn test_progress_phase_weight() {
assert!(ProgressPhase::Encoding.weight() > ProgressPhase::Resampling.weight());
assert!(ProgressPhase::Decoding.weight() > ProgressPhase::Encoding.weight());
assert_eq!(ProgressPhase::Complete.weight(), 0);
}
#[test]
fn test_progress_phase_weights_sum() {
let total: u32 = [
ProgressPhase::LoadingModel,
ProgressPhase::Resampling,
ProgressPhase::MelSpectrogram,
ProgressPhase::Encoding,
ProgressPhase::Decoding,
ProgressPhase::Postprocessing,
]
.iter()
.map(|p| p.weight())
.sum();
assert_eq!(total, 100); }
#[test]
fn test_worker_progress_new() {
let progress = WorkerProgress::new();
assert_eq!(progress.phase(), ProgressPhase::LoadingModel);
assert!((progress.phase_progress() - 0.0).abs() < f32::EPSILON);
assert!((progress.overall_progress() - 0.0).abs() < f32::EPSILON);
}
#[test]
fn test_worker_progress_update() {
let mut progress = WorkerProgress::new();
progress.update(ProgressPhase::Encoding, 50.0, "Processing...");
assert_eq!(progress.phase(), ProgressPhase::Encoding);
assert!((progress.phase_progress() - 50.0).abs() < f32::EPSILON);
assert_eq!(progress.message(), "Processing...");
assert!(progress.overall_progress() > 0.0);
}
#[test]
fn test_worker_progress_next_phase() {
let mut progress = WorkerProgress::new();
progress.next_phase(ProgressPhase::Resampling);
assert_eq!(progress.phase(), ProgressPhase::Resampling);
assert!((progress.phase_progress() - 0.0).abs() < f32::EPSILON);
progress.next_phase(ProgressPhase::MelSpectrogram);
assert_eq!(progress.phase(), ProgressPhase::MelSpectrogram);
}
#[test]
fn test_worker_progress_complete() {
let mut progress = WorkerProgress::new();
progress.complete();
assert_eq!(progress.phase(), ProgressPhase::Complete);
assert!((progress.phase_progress() - 100.0).abs() < f32::EPSILON);
assert!((progress.overall_progress() - 100.0).abs() < f32::EPSILON);
assert_eq!(progress.message(), "Complete");
}
#[test]
fn test_worker_progress_overall_increases() {
let mut progress = WorkerProgress::new();
let initial = progress.overall_progress();
progress.next_phase(ProgressPhase::Encoding);
let after_encoding = progress.overall_progress();
assert!(after_encoding > initial);
progress.next_phase(ProgressPhase::Decoding);
let after_decoding = progress.overall_progress();
assert!(after_decoding > after_encoding);
}
#[test]
fn test_worker_progress_clamps_phase_progress() {
let mut progress = WorkerProgress::new();
progress.update(ProgressPhase::Encoding, 150.0, "Over 100");
assert!((progress.phase_progress() - 100.0).abs() < f32::EPSILON);
progress.update(ProgressPhase::Encoding, -50.0, "Under 0");
assert!((progress.phase_progress() - 0.0).abs() < f32::EPSILON);
}
#[test]
fn test_worker_config_default() {
let config = WorkerConfig::default();
assert_eq!(config.model_type(), "tiny");
assert!(config.enable_progress());
assert_eq!(config.progress_interval_ms(), 100);
}
#[test]
fn test_worker_config_setters() {
let mut config = WorkerConfig::new();
config.set_model_type("base");
assert_eq!(config.model_type(), "base");
config.set_enable_progress(false);
assert!(!config.enable_progress());
config.set_progress_interval_ms(200);
assert_eq!(config.progress_interval_ms(), 200);
}
}