Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 51 additions & 8 deletions src-tauri/src/dictation/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -88,13 +88,18 @@ impl Inner {
model::path_in(&self.dir)
}

fn usable_model(&self) -> Option<PathBuf> {
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,
Expand Down Expand Up @@ -142,10 +147,42 @@ pub fn dictation_status(state: State<Dictation>) -> Status {

#[tauri::command]
pub fn dictation_download_model(state: State<Dictation>) -> 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<Inner>) -> 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());
}
Expand All @@ -170,6 +207,7 @@ pub fn dictation_download_model(state: State<Dictation>) -> Result<(), String> {
crate::log::info("dictation.downloaded")
.elapsed(started)
.emit();
remove_previous_model(&inner);
inner.set_phase(Phase::Idle)
}

Expand All @@ -193,9 +231,10 @@ pub fn dictation_cancel_download(state: State<Dictation>) {
#[tauri::command]
pub fn dictation_delete_model(state: State<Dictation>) -> 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();
Expand All @@ -208,7 +247,7 @@ pub fn dictation_start(state: State<Dictation>) -> 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)?;
Expand Down Expand Up @@ -260,7 +299,11 @@ fn finish_recording(inner: &Arc<Inner>, language: Option<String>) -> 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.");
}
Expand Down
80 changes: 73 additions & 7 deletions src-tauri/src/dictation/model.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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<PathBuf> {
let current = path_in(dir);
if is_installed(&current) {
return Some(current);
}
previous_installed(dir).then(|| previous_path_in(dir))
}

#[derive(Debug)]
pub enum DownloadError {
Cancelled,
Expand Down Expand Up @@ -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);
}
}
4 changes: 3 additions & 1 deletion src-tauri/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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());

Expand Down
Loading