diff --git a/src-tauri/src/dictation/mod.rs b/src-tauri/src/dictation/mod.rs index c39e920..028f219 100644 --- a/src-tauri/src/dictation/mod.rs +++ b/src-tauri/src/dictation/mod.rs @@ -88,13 +88,18 @@ impl Inner { model::path_in(&self.dir) } + fn usable_model(&self) -> Option { + model::usable_in(&self.dir) + } + fn status(&self) -> Status { - let path = self.model_path(); + let usable = self.usable_model(); + let path = usable.clone().unwrap_or_else(|| self.model_path()); Status { available: cfg!(feature = "dictation"), phase: *self.phase.lock().unwrap(), - model_installed: model::is_installed(&path), + model_installed: usable.is_some(), model_path: Some(path.to_string_lossy().into_owned()), model_bytes: model::BYTES, model_memory_bytes: model::MEMORY_BYTES, @@ -142,10 +147,42 @@ pub fn dictation_status(state: State) -> Status { #[tauri::command] pub fn dictation_download_model(state: State) -> Result<(), String> { + download_model(state.0.clone()) +} + +impl Dictation { + pub fn upgrade_model(&self) { + let inner = &self.0; + if !cfg!(feature = "dictation") || !model::previous_installed(&inner.dir) { + return; + } + if model::is_installed(&inner.model_path()) { + remove_previous_model(inner); + return; + } + crate::log::info("dictation.model_upgrade") + .field("from", model::previous::FILE) + .field("to", model::FILE) + .emit(); + let _ = download_model(inner.clone()); + } +} + +fn remove_previous_model(inner: &Inner) { + let path = model::previous_path_in(&inner.dir); + if let Err(e) = std::fs::remove_file(&path) { + if e.kind() != std::io::ErrorKind::NotFound { + crate::log::warn("dictation.previous_model_kept") + .field("error", e.to_string()) + .emit(); + } + } +} + +fn download_model(inner: Arc) -> Result<(), String> { if !cfg!(feature = "dictation") { return Err(UNAVAILABLE.into()); } - let inner = state.0.clone(); if *inner.phase.lock().unwrap() != Phase::Idle { return Err("Dictation is already busy".into()); } @@ -170,6 +207,7 @@ pub fn dictation_download_model(state: State) -> Result<(), String> { crate::log::info("dictation.downloaded") .elapsed(started) .emit(); + remove_previous_model(&inner); inner.set_phase(Phase::Idle) } @@ -193,9 +231,10 @@ pub fn dictation_cancel_download(state: State) { #[tauri::command] pub fn dictation_delete_model(state: State) -> Result<(), String> { let inner = &state.0; - let path = inner.model_path(); - if path.exists() { - std::fs::remove_file(&path).map_err(|e| inner.fail(format!("Couldn't delete: {e}")))?; + for path in [inner.model_path(), model::previous_path_in(&inner.dir)] { + if path.exists() { + std::fs::remove_file(&path).map_err(|e| inner.fail(format!("Couldn't delete: {e}")))?; + } } inner.clear_error(); inner.publish(); @@ -208,7 +247,7 @@ pub fn dictation_start(state: State) -> Result<(), String> { if *inner.phase.lock().unwrap() != Phase::Idle { return Err("Dictation is already busy".into()); } - if !model::is_installed(&inner.model_path()) { + if inner.usable_model().is_none() { return Err(inner.fail("The dictation model isn't downloaded yet.")); } start_recording(&inner)?; @@ -260,7 +299,11 @@ fn finish_recording(inner: &Arc, language: Option) -> Result<(), return; } let started = std::time::Instant::now(); - match engine::transcribe(&inner.model_path(), &audio, language.as_deref()) { + let Some(model) = inner.usable_model() else { + inner.fail("The dictation model isn't downloaded yet."); + return; + }; + match engine::transcribe(&model, &audio, language.as_deref()) { Ok(text) if text.is_empty() => { inner.fail("Nothing was picked up. Check the microphone."); } diff --git a/src-tauri/src/dictation/model.rs b/src-tauri/src/dictation/model.rs index 9b94cc1..78176a2 100644 --- a/src-tauri/src/dictation/model.rs +++ b/src-tauri/src/dictation/model.rs @@ -2,21 +2,31 @@ use std::fs; use std::path::{Path, PathBuf}; use std::sync::atomic::AtomicBool; -pub const FILE: &str = "ggml-small.bin"; +pub const FILE: &str = "ggml-small-q5_1.bin"; -pub const URL: &str = "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-small.bin"; +pub const URL: &str = + "https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-small-q5_1.bin"; -pub const BYTES: u64 = 487_601_967; +pub const BYTES: u64 = 190_085_487; -pub const SHA256: &str = "1be3a9b2063867b937e64e2ec7483364a79917e157fa98c5d94b5c1fffea987b"; +pub const SHA256: &str = "ae85e4a935d7a567bd102fe55afc16bb595bdb618e11b2fc7591bc08120411bb"; -pub const MEMORY_BYTES: u64 = 900_000_000; +pub const MEMORY_BYTES: u64 = 700_000_000; + +pub mod previous { + pub const FILE: &str = "ggml-small.bin"; + pub const BYTES: u64 = 487_601_967; +} pub const LABEL: &str = "Whisper small"; pub const PARAMETERS: &str = "244M"; fn plausible(len: u64) -> bool { - len >= BYTES / 2 + plausible_for(len, BYTES) +} + +fn plausible_for(len: u64, expected: u64) -> bool { + len >= expected / 2 } pub fn is_installed(path: &Path) -> bool { @@ -25,6 +35,24 @@ pub fn is_installed(path: &Path) -> bool { .unwrap_or(false) } +pub fn previous_path_in(dir: &Path) -> PathBuf { + dir.join(previous::FILE) +} + +pub fn previous_installed(dir: &Path) -> bool { + fs::metadata(previous_path_in(dir)) + .map(|m| plausible_for(m.len(), previous::BYTES)) + .unwrap_or(false) +} + +pub fn usable_in(dir: &Path) -> Option { + let current = path_in(dir); + if is_installed(¤t) { + return Some(current); + } + previous_installed(dir).then(|| previous_path_in(dir)) +} + #[derive(Debug)] pub enum DownloadError { Cancelled, @@ -149,7 +177,45 @@ mod tests { fn the_model_path_is_the_whisper_cpp_name() { assert_eq!( path_in(Path::new("/models")), - Path::new("/models/ggml-small.bin") + Path::new("/models/ggml-small-q5_1.bin") ); } + + fn model_file(dir: &Path, name: &str, len: u64) { + fs::File::create(dir.join(name)) + .unwrap() + .set_len(len) + .unwrap(); + } + + #[test] + fn the_current_model_is_used_when_it_is_installed() { + let dir = crate::testing::TempDir::new("model"); + model_file(dir.path(), FILE, BYTES); + model_file(dir.path(), previous::FILE, previous::BYTES); + assert_eq!(usable_in(dir.path()), Some(path_in(dir.path()))); + } + + #[test] + fn the_previous_model_keeps_dictation_working_until_it_is_replaced() { + let dir = crate::testing::TempDir::new("model"); + model_file(dir.path(), previous::FILE, previous::BYTES); + assert_eq!(usable_in(dir.path()), Some(previous_path_in(dir.path()))); + assert!(!is_installed(&path_in(dir.path()))); + } + + #[test] + fn a_half_downloaded_current_model_falls_back_to_the_previous_one() { + let dir = crate::testing::TempDir::new("model"); + model_file(dir.path(), FILE, BYTES / 4); + model_file(dir.path(), previous::FILE, previous::BYTES); + assert_eq!(usable_in(dir.path()), Some(previous_path_in(dir.path()))); + } + + #[test] + fn a_truncated_previous_model_is_not_used() { + let dir = crate::testing::TempDir::new("model"); + model_file(dir.path(), previous::FILE, previous::BYTES / 4); + assert_eq!(usable_in(dir.path()), None); + } } diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index bd3f480..8edee4a 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -284,7 +284,9 @@ pub fn run() { app.manage(watch::Watch::new(app.handle().clone())); - app.manage(dictation::Dictation::new(app.handle().clone())); + let dictation = dictation::Dictation::new(app.handle().clone()); + dictation.upgrade_model(); + app.manage(dictation); app.manage(updates::Updates::default());