#![cfg_attr(not(feature = "local-stt"), allow(dead_code))]
use std::future::Future;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, Ordering};
#[cfg(feature = "local-stt")]
mod engine;
#[cfg(feature = "local-stt")]
mod wav;
pub const ENABLED: bool = cfg!(feature = "local-stt");
pub const DEFAULT_MODEL: &str = "base.en";
pub const MODEL_DIR_ENV: &str = "MOBUX_STT_MODEL_DIR";
pub const ASSET_BASE_URL_ENV: &str = "MOBUX_STT_ASSET_BASE_URL";
pub fn default_asset_base_url() -> String {
format!(
"https://github.com/mvhenten/mobux/releases/download/v{}",
env!("CARGO_PKG_VERSION")
)
}
pub fn asset_model_prefix(model: &str) -> String {
format!("stt-models/{model}/")
}
const MODEL_LOCK_JSON: &str = include_str!("local_stt/model.lock.json");
pub use crate::release_asset::LockedFile;
#[derive(Debug, Clone, serde::Deserialize)]
pub struct LockedModel {
pub id: String,
pub files: std::collections::BTreeMap<String, LockedFile>,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct ModelLock {
#[serde(rename = "default")]
pub default_model: String,
pub vendored: String,
pub models: Vec<LockedModel>,
}
pub fn model_lock() -> &'static ModelLock {
static LOCK: std::sync::OnceLock<ModelLock> = std::sync::OnceLock::new();
LOCK.get_or_init(|| {
serde_json::from_str(MODEL_LOCK_JSON).expect("model.lock.json is built into the binary")
})
}
pub fn locked_model(model: &str) -> Option<&'static LockedModel> {
let model = model.trim();
model_lock().models.iter().find(|m| m.id == model)
}
pub fn model_files(model: &str) -> Vec<&'static str> {
locked_model(model)
.map(|m| m.files.keys().map(String::as_str).collect())
.unwrap_or_default()
}
pub fn model_ids() -> Vec<String> {
model_lock().models.iter().map(|m| m.id.clone()).collect()
}
pub fn is_known_model(model: &str) -> bool {
locked_model(model).is_some()
}
pub fn resolve_model(configured: &str) -> &'static str {
locked_model(configured)
.map(|m| m.id.as_str())
.unwrap_or(model_lock().default_model.as_str())
}
pub fn cache_dir(data_dir: &Path) -> PathBuf {
data_dir.join("stt-models")
}
pub fn model_dir(data_dir: &Path, model: &str) -> PathBuf {
cache_dir(data_dir).join(model)
}
pub fn files_present(dir: &Path, model: &str) -> bool {
let files = model_files(model);
!files.is_empty() && files.iter().all(|f| dir.join(f).is_file())
}
pub fn model_files_present(data_dir: &Path, model: &str) -> bool {
files_present(&model_dir(data_dir, model), model)
}
pub use crate::release_asset::release_asset_name;
pub fn asset_for(model: &str) -> Option<String> {
let model = resolve_model(model);
if model == model_lock().vendored {
return release_asset_name().map(str::to_string);
}
Some(format!("mobux-stt-{model}.tar.gz"))
}
pub fn asset_base_url() -> String {
std::env::var(ASSET_BASE_URL_ENV)
.ok()
.filter(|u| !u.is_empty())
.unwrap_or_else(default_asset_base_url)
}
#[cfg(all(feature = "local-stt", target_arch = "aarch64"))]
pub fn unsupported_cpu() -> Option<String> {
if std::arch::is_aarch64_feature_detected!("fp16") {
return None;
}
Some(CPU_UNSUPPORTED_MESSAGE.to_string())
}
#[cfg(not(all(feature = "local-stt", target_arch = "aarch64")))]
pub fn unsupported_cpu() -> Option<String> {
None
}
#[cfg(any(all(feature = "local-stt", target_arch = "aarch64"), test))]
pub const CPU_UNSUPPORTED_MESSAGE: &str =
"This CPU has no ARMv8.2 half-precision support (FEAT_FP16), which the in-process speech engine needs — a Raspberry Pi 4 and other ARMv8.0 cores do not have it. Point the provider at an OpenAI-compatible endpoint instead.";
#[derive(Debug, Clone, PartialEq)]
pub enum Phase {
Disabled,
UnsupportedCpu(String),
NotDownloaded,
Verifying,
Downloading {
file: String,
downloaded: u64,
total: u64,
},
Loading,
Ready,
Failed(String),
}
impl Phase {
pub fn state(&self) -> &'static str {
match self {
Self::Disabled | Self::UnsupportedCpu(_) => "unsupported",
Self::NotDownloaded => "not_installed",
Self::Downloading { .. } | Self::Verifying | Self::Loading => "warming",
Self::Ready => "ready",
Self::Failed(_) => "failed",
}
}
pub fn message(&self) -> String {
match self {
Self::Disabled => UNSUPPORTED_MESSAGE.to_string(),
Self::UnsupportedCpu(why) => why.clone(),
Self::NotDownloaded => "The speech model is not on this host yet.".to_string(),
Self::Downloading {
file,
downloaded,
total,
} => match percent(*downloaded, *total) {
Some(pct) => format!("Downloading the speech model ({file}) — {pct}%."),
None => format!("Downloading the speech model ({file})."),
},
Self::Verifying => "Checking the speech model against its recorded hashes.".to_string(),
Self::Loading => "Loading the speech model into memory.".to_string(),
Self::Ready => "The speech model is loaded.".to_string(),
Self::Failed(err) => format!("The speech model could not be prepared: {err}"),
}
}
}
pub const UNSUPPORTED_MESSAGE: &str =
"This build has no in-process speech engine. Reinstall with `cargo install mobux --locked --features local-stt`, or point the provider at an OpenAI-compatible endpoint.";
fn percent(downloaded: u64, total: u64) -> Option<u64> {
if total == 0 {
return None;
}
Some((downloaded.saturating_mul(100) / total).min(100))
}
#[cfg(feature = "local-stt")]
pub fn phase(data_dir: &Path, model: &str) -> Phase {
if let Some(why) = unsupported_cpu() {
return Phase::UnsupportedCpu(why);
}
engine::phase(data_dir, model)
}
#[cfg(not(feature = "local-stt"))]
pub fn phase(_data_dir: &Path, _model: &str) -> Phase {
Phase::Disabled
}
#[cfg(feature = "local-stt")]
pub async fn ensure_ready(data_dir: PathBuf, model: String) -> Result<(), String> {
if let Some(why) = unsupported_cpu() {
return Err(why);
}
engine::ensure_ready(data_dir, model).await
}
#[cfg(not(feature = "local-stt"))]
pub async fn ensure_ready(_data_dir: PathBuf, _model: String) -> Result<(), String> {
Err(UNSUPPORTED_MESSAGE.to_string())
}
pub struct SingleFlight {
busy: AtomicBool,
}
impl SingleFlight {
pub const fn new() -> Self {
Self {
busy: AtomicBool::new(false),
}
}
pub fn spawn<F>(&'static self, work: F) -> bool
where
F: Future<Output = ()> + Send + 'static,
{
if self
.busy
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
return false;
}
tokio::spawn(async move {
let _release = Release(&self.busy);
work.await;
});
true
}
}
impl Default for SingleFlight {
fn default() -> Self {
Self::new()
}
}
struct Release(&'static AtomicBool);
impl Drop for Release {
fn drop(&mut self) {
self.0.store(false, Ordering::Release);
}
}
pub fn prepare_in_background(data_dir: PathBuf, model: String) -> bool {
static PREPARING: SingleFlight = SingleFlight::new();
PREPARING.spawn(async move {
if let Err(e) = ensure_ready(data_dir, model).await {
eprintln!("[stt] preparing the speech model failed: {e}");
}
})
}
#[cfg(feature = "local-stt")]
pub async fn transcribe(data_dir: PathBuf, model: String, wav: Vec<u8>) -> Result<String, String> {
if let Some(why) = unsupported_cpu() {
return Err(why);
}
engine::transcribe(data_dir, model, wav).await
}
#[cfg(not(feature = "local-stt"))]
pub async fn transcribe(
_data_dir: PathBuf,
_model: String,
_wav: Vec<u8>,
) -> Result<String, String> {
Err(UNSUPPORTED_MESSAGE.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn polls_during_a_preparation_start_it_once() {
static FLIGHT: SingleFlight = SingleFlight::new();
static RUNS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
let (finish, finished) = tokio::sync::oneshot::channel::<()>();
let (done, wait_done) = tokio::sync::oneshot::channel::<()>();
let started = FLIGHT.spawn(async move {
RUNS.fetch_add(1, Ordering::SeqCst);
let _ = finished.await;
let _ = done.send(());
});
assert!(started);
for _ in 0..5 {
assert!(
!FLIGHT.spawn(async {
RUNS.fetch_add(1, Ordering::SeqCst);
}),
"a poll queued a second preparation"
);
}
finish.send(()).unwrap();
wait_done.await.unwrap();
tokio::task::yield_now().await;
assert_eq!(RUNS.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_finished_preparation_frees_the_slot() {
static FLIGHT: SingleFlight = SingleFlight::new();
let (done, wait_done) = tokio::sync::oneshot::channel::<()>();
assert!(FLIGHT.spawn(async move {
let _ = done.send(());
}));
wait_done.await.unwrap();
for _ in 0..100 {
if !FLIGHT.busy.load(Ordering::Acquire) {
break;
}
tokio::task::yield_now().await;
}
assert!(
FLIGHT.spawn(async {}),
"the slot stayed taken after the run ended"
);
}
#[test]
fn the_default_is_the_vendored_model_and_the_card_offers_it_first() {
let lock = model_lock();
assert_eq!(lock.default_model, DEFAULT_MODEL);
assert_eq!(lock.vendored, DEFAULT_MODEL);
assert!(is_known_model(DEFAULT_MODEL));
assert_eq!(model_ids().first().map(String::as_str), Some(DEFAULT_MODEL));
}
#[test]
fn the_catalog_is_the_three_english_checkpoints() {
assert_eq!(model_ids(), vec!["base.en", "tiny.en", "small.en"]);
}
#[test]
fn every_locked_checkpoint_pins_three_f16_files_with_real_hashes() {
const F32_WOULD_EXCEED: u64 = 600 * 1024 * 1024;
for model in model_lock().models.iter() {
assert_eq!(
model_files(&model.id),
vec!["config.json", "model.safetensors", "tokenizer.json"],
"{}",
model.id
);
for (name, file) in &model.files {
assert_eq!(file.sha256.len(), 64, "{}/{name}", model.id);
assert!(
file.sha256.chars().all(|c| c.is_ascii_hexdigit()),
"{}/{name}",
model.id
);
assert!(file.bytes > 0, "{}/{name}", model.id);
}
assert!(
model.files["model.safetensors"].bytes < F32_WOULD_EXCEED,
"{} is {} bytes — not converted to f16?",
model.id,
model.files["model.safetensors"].bytes
);
}
}
#[test]
fn the_vendored_model_comes_from_the_platform_tarball_and_the_rest_from_their_own() {
assert_eq!(
asset_for(DEFAULT_MODEL),
release_asset_name().map(str::to_string)
);
assert_eq!(
asset_for("tiny.en").as_deref(),
Some("mobux-stt-tiny.en.tar.gz")
);
assert_eq!(
asset_for("small.en").as_deref(),
Some("mobux-stt-small.en.tar.gz")
);
}
#[test]
fn the_asset_url_is_pinned_to_this_builds_own_release() {
let url = default_asset_base_url();
assert!(
url.ends_with(&format!(
"/releases/download/v{}",
env!("CARGO_PKG_VERSION")
)),
"{url}"
);
assert!(!url.contains("latest"), "{url}");
}
#[test]
fn a_cpu_that_cannot_run_the_engine_reads_as_unsupported_with_a_reason() {
let phase = Phase::UnsupportedCpu(CPU_UNSUPPORTED_MESSAGE.to_string());
assert_eq!(phase.state(), "unsupported");
assert!(phase.message().contains("FEAT_FP16"));
assert!(phase.message().contains("OpenAI-compatible"));
}
#[test]
fn every_catalog_model_has_an_asset_and_a_path_inside_it() {
assert!(asset_base_url().starts_with("http"));
for id in model_ids() {
assert!(
asset_for(&id).is_some_and(|a| a.ends_with(".tar.gz")),
"{id}"
);
assert_eq!(asset_model_prefix(&id), format!("stt-models/{id}/"));
}
}
#[test]
fn a_model_the_engine_cannot_run_falls_back_to_the_default() {
assert_eq!(resolve_model("Systran/faster-whisper-small"), DEFAULT_MODEL);
assert_eq!(resolve_model(""), DEFAULT_MODEL);
assert_eq!(resolve_model("whisper-1"), DEFAULT_MODEL);
}
#[test]
fn a_model_in_the_catalog_is_kept() {
assert_eq!(resolve_model(DEFAULT_MODEL), DEFAULT_MODEL);
assert_eq!(resolve_model(" small.en "), "small.en");
assert_eq!(resolve_model("tiny.en"), "tiny.en");
}
fn write_model_files(dir: &Path, model: &str) {
std::fs::create_dir_all(dir).unwrap();
for name in model_files(model) {
std::fs::write(dir.join(name), b"x").unwrap();
}
}
#[test]
fn model_files_are_reported_missing_until_all_three_exist() {
let dir = tempfile::tempdir().unwrap();
assert!(!model_files_present(dir.path(), DEFAULT_MODEL));
let model = model_dir(dir.path(), DEFAULT_MODEL);
std::fs::create_dir_all(&model).unwrap();
std::fs::write(model.join("config.json"), b"{}").unwrap();
std::fs::write(model.join("tokenizer.json"), b"{}").unwrap();
assert!(!model_files_present(dir.path(), DEFAULT_MODEL));
std::fs::write(model.join("model.safetensors"), b"x").unwrap();
assert!(model_files_present(dir.path(), DEFAULT_MODEL));
}
#[test]
fn a_directory_counts_as_present_only_once_it_holds_every_file() {
let dir = tempfile::tempdir().unwrap();
assert!(!files_present(dir.path(), DEFAULT_MODEL));
write_model_files(dir.path(), DEFAULT_MODEL);
assert!(files_present(dir.path(), DEFAULT_MODEL));
}
#[test]
fn each_model_caches_in_its_own_directory() {
let dir = tempfile::tempdir().unwrap();
write_model_files(&model_dir(dir.path(), "small.en"), "small.en");
assert!(model_files_present(dir.path(), "small.en"));
assert!(!model_files_present(dir.path(), DEFAULT_MODEL));
assert!(!model_files_present(dir.path(), "tiny.en"));
}
#[test]
fn a_model_outside_the_catalog_has_no_files_to_look_for() {
let dir = tempfile::tempdir().unwrap();
assert!(model_files("medium.en").is_empty());
assert!(!files_present(dir.path(), "medium.en"));
}
#[test]
fn phases_render_a_state_word_and_a_sentence() {
assert_eq!(Phase::Ready.state(), "ready");
assert_eq!(Phase::NotDownloaded.state(), "not_installed");
assert_eq!(Phase::Loading.state(), "warming");
assert_eq!(Phase::Verifying.state(), "warming");
assert_eq!(
Phase::Downloading {
file: "model.safetensors".to_string(),
downloaded: 50,
total: 200,
}
.state(),
"warming"
);
assert_eq!(Phase::Failed("boom".to_string()).state(), "failed");
assert_eq!(Phase::Disabled.state(), "unsupported");
assert_eq!(
Phase::UnsupportedCpu(CPU_UNSUPPORTED_MESSAGE.to_string()).state(),
"unsupported"
);
let msg = Phase::Downloading {
file: "model.safetensors".to_string(),
downloaded: 50,
total: 200,
}
.message();
assert!(msg.contains("25%"), "{msg}");
assert!(Phase::Disabled.message().contains("--features local-stt"));
}
#[test]
fn an_unknown_download_total_reports_no_percentage() {
let msg = Phase::Downloading {
file: "model.safetensors".to_string(),
downloaded: 4096,
total: 0,
}
.message();
assert!(!msg.contains('%'), "{msg}");
}
}