This commit is contained in:
iamdoubz
2026-06-30 23:53:33 -05:00
parent 1d3c8c3749
commit ba48f0afc2
161 changed files with 30574 additions and 54 deletions
+229 -29
View File
@@ -7,11 +7,14 @@
use crate::audio::{AudioCapture, WasapiCapture};
use crate::error::WaResult;
use crate::hardware::{HardwareDetector, WinHardwareDetector};
use crate::models::*;
use crate::notes::NotesRenderer;
use crate::paths::{meeting_dir, settings_path, wa_root, whisper_model_path};
use crate::paths::{meeting_dir, settings_path, wa_root, whisper_model_file};
use crate::storage::{FinalizeMeeting, Meeting, NewMeeting};
use crate::transcription::{run_streaming_worker, Transcriber, WhisperTranscriber};
use crate::transcription::{
models as model_catalog, run_streaming_worker, Transcriber, WhisperTranscriber,
};
use crate::{error::WaError, AppState, RecordingSession};
use serde::Deserialize;
use std::path::PathBuf;
@@ -39,6 +42,7 @@ fn default_settings() -> Settings {
llm_endpoint: "http://localhost:11434".into(),
llm_model: "llama3".into(),
preferred_backend: "auto".into(),
whisper_model: crate::paths::DEFAULT_WHISPER_MODEL.to_string(),
low_overhead: false,
default_record: false,
consent_acknowledged: false,
@@ -65,6 +69,35 @@ fn save_settings(settings: &Settings) -> Result<(), WaError> {
std::fs::write(path, json).map_err(|e| WaError::new("settings", e.to_string()))
}
/// Resolves which model file a recording should use: the "low overhead" preset
/// (T3.9) always takes the smallest catalog model over whatever is configured,
/// since it's optimizing for overhead, not accuracy.
fn model_id_for(settings: &Settings) -> String {
if settings.low_overhead {
model_catalog::smallest_id().to_string()
} else {
settings.whisper_model.clone()
}
}
/// Picks the backend to transcribe with: "low overhead" always forces CPU;
/// otherwise resolve the user's preferred backend (or auto-detect) against
/// what's actually available (T3.2, T3.5).
fn backend_for(settings: &Settings) -> BackendId {
if settings.low_overhead {
return BackendId::Cpu;
}
let preferred = match settings.preferred_backend.as_str() {
"npu" => Some(BackendId::Npu),
"nvidia" => Some(BackendId::Nvidia),
"amd" => Some(BackendId::Amd),
"intel" => Some(BackendId::Intel),
"cpu" => Some(BackendId::Cpu),
_ => None, // "auto"
};
WinHardwareDetector.best(preferred).id
}
// ---- Recording lifecycle (Phase 1) ----
#[tauri::command]
@@ -88,12 +121,15 @@ pub async fn start_recording(
));
}
let model_path = whisper_model_path();
let settings = load_settings();
let model_id = model_id_for(&settings);
let backend = backend_for(&settings);
let model_path = whisper_model_file(&model_id);
if !model_path.exists() {
return Err(WaError::new(
"transcription",
format!(
"whisper model not found at {}; run `npm run download-models` first",
"whisper model '{model_id}' not found at {}; download it from Settings first",
model_path.display()
),
));
@@ -119,24 +155,50 @@ pub async fn start_recording(
let segments: Arc<StdMutex<Vec<TranscriptSegment>>> = Arc::new(StdMutex::new(Vec::new()));
let segments_for_worker = segments.clone();
let active_backend: Arc<StdMutex<BackendId>> = Arc::new(StdMutex::new(backend));
let active_backend_for_worker = active_backend.clone();
let app_for_worker = app.clone();
let meeting_id_for_worker = meeting_id.clone();
let transcription_worker = std::thread::Builder::new()
.name("wa-transcription".into())
.spawn(
move || match WhisperTranscriber::load(&model_path, BackendId::Cpu) {
Ok(transcriber) => run_streaming_worker(&transcriber, frame_rx, |segment| {
if let Ok(mut buf) = segments_for_worker.lock() {
buf.push(segment.clone());
.spawn(move || {
// Graceful fallback (T3.5, FR-HW-4): a GPU/NPU load failure (driver
// issue, OOM, unsupported adapter) falls back to CPU rather than
// losing the meeting's transcript entirely.
let transcriber = match WhisperTranscriber::load(&model_path, backend) {
Ok(t) => t,
Err(e) if backend != BackendId::Cpu => {
tracing::warn!("backend {backend:?} failed to load ({e}); falling back to CPU");
if let Ok(mut b) = active_backend_for_worker.lock() {
*b = BackendId::Cpu;
}
let _ = app_for_worker.emit(
"hardware://changed",
serde_json::json!({ "active": BackendId::Cpu.as_str(), "reason": e.to_string() }),
);
match WhisperTranscriber::load(&model_path, BackendId::Cpu) {
Ok(t) => t,
Err(e) => {
tracing::error!("failed to load whisper model on CPU fallback: {e}");
return;
}
}
}
Err(e) => {
tracing::error!("failed to load whisper model: {e}");
return;
}
};
run_streaming_worker(&transcriber, frame_rx, |segment| {
if let Ok(mut buf) = segments_for_worker.lock() {
buf.push(segment.clone());
}
let _ = app_for_worker.emit(
"transcript://segment",
serde_json::json!({ "meetingId": meeting_id_for_worker, "segment": segment }),
);
}),
Err(e) => tracing::error!("failed to load whisper model: {e}"),
},
)
});
})
.map_err(|e| WaError::new("transcription", e.to_string()))?;
*guard = Some(RecordingSession {
@@ -147,6 +209,8 @@ pub async fn start_recording(
started_at: std::time::Instant::now(),
transcription_worker,
segments,
active_backend,
model_id,
});
drop(guard);
@@ -203,10 +267,11 @@ pub async fn stop_recording(
display_name: None,
participant_id: None,
}];
let model_used = whisper_model_path()
.file_stem()
.and_then(|s| s.to_str())
.map(str::to_string);
let backend_used = session
.active_backend
.lock()
.map(|b| b.as_str().to_string())
.unwrap_or_else(|_| BackendId::Cpu.as_str().to_string());
// T2.10: persist transcript.json (via finalize_meeting) and notes.md
// *before* touching the working WAV, so a crash here still leaves a
@@ -221,8 +286,8 @@ pub async fn stop_recording(
duration_secs: (summary.duration_ms / 1000) as i64,
recorded: session.retention,
language: None,
backend_used: Some(BackendId::Cpu.as_str().to_string()),
model_used,
backend_used: Some(backend_used),
model_used: Some(session.model_id.clone()),
},
)
.await
@@ -330,11 +395,148 @@ pub async fn acknowledge_recording_consent() -> WaResult<()> {
save_settings(&settings)
}
// ---- Hardware (Phase 3) ----
// ---- Hardware + models (Phase 3) ----
#[tauri::command]
pub async fn hardware_status() -> WaResult<Vec<BackendInfo>> {
todo!("Phase 3 — hardware_status")
pub async fn hardware_status() -> WaResult<serde_json::Value> {
let settings = load_settings();
let backends = WinHardwareDetector.detect();
let active = backend_for(&settings);
let model_id = model_id_for(&settings);
Ok(serde_json::json!({
"backends": backends,
"active": active,
"modelSize": model_id,
// ponytail: no real-time factor measurement harness yet (needs a
// timed sample transcription); 1.0 stands in until T3.6 wires one up.
"estRtf": 1.0,
}))
}
#[derive(Deserialize)]
pub struct SetPreferredBackendArgs {
pub backend: String, // "auto"|"npu"|"nvidia"|"amd"|"intel"|"cpu"
}
#[tauri::command]
pub async fn set_preferred_backend(args: SetPreferredBackendArgs) -> WaResult<()> {
let mut settings = load_settings();
settings.preferred_backend = args.backend;
save_settings(&settings)
}
#[tauri::command]
pub async fn list_models() -> WaResult<Vec<ModelInfo>> {
let settings = load_settings();
Ok(model_catalog::list(&model_id_for(&settings)))
}
#[derive(Deserialize)]
pub struct DownloadModelArgs {
pub kind: String, // "whisper" — diar-seg/diar-emb land in Phase 4 (T4.7)
pub id: String,
}
#[tauri::command]
pub async fn download_model(app: AppHandle, args: DownloadModelArgs) -> WaResult<()> {
if args.kind != "whisper" {
return Err(WaError::new(
"model",
format!("model kind '{}' is not available yet", args.kind),
));
}
let id = args.id;
let id_for_progress = id.clone();
model_catalog::download(&id, move |received, total| {
let _ = app.emit(
"model://progress",
serde_json::json!({ "id": id_for_progress, "receivedBytes": received, "totalBytes": total }),
);
})
.await
.map_err(|e| WaError::new("model", e.to_string()))
}
#[tauri::command]
pub async fn remove_model(id: String) -> WaResult<()> {
let settings = load_settings();
model_catalog::remove(&id, &model_id_for(&settings))
.map_err(|e| WaError::new("model", e.to_string()))
}
/// Batch re-transcribe a finished meeting with a different (typically larger)
/// model (T3.8, FR-TRX-3). Only works if the meeting's audio was retained.
#[tauri::command]
pub async fn reprocess_transcript(
app: AppHandle,
state: State<'_, AppState>,
meeting_id: MeetingId,
model: String,
) -> WaResult<()> {
let model_path = whisper_model_file(&model);
if !model_path.exists() {
return Err(WaError::new(
"model",
format!("model '{model}' is not installed"),
));
}
let wav_path = meeting_dir(&meeting_id).join("audio.wav");
if !wav_path.exists() {
return Err(WaError::new(
"transcription",
"no retained audio.wav to reprocess — enable recording retention for this meeting",
));
}
let settings = load_settings();
let backend = backend_for(&settings);
let segments = tauri::async_runtime::spawn_blocking({
let wav_path = wav_path.clone();
move || {
WhisperTranscriber::load(&model_path, backend)
.or_else(|_| WhisperTranscriber::load(&model_path, BackendId::Cpu))
.and_then(|transcriber| transcriber.transcribe_file(&wav_path))
}
})
.await
.map_err(|e| WaError::new("transcription", e.to_string()))?
.map_err(|e| WaError::new("transcription", e.to_string()))?;
let meeting = state
.store
.get_meeting(&meeting_id)
.await
.map_err(|e| WaError::new("storage", e.to_string()))?;
let duration_secs = segments
.last()
.map(|s| (s.end_ms / 1000) as i64)
.unwrap_or(meeting.duration_secs.unwrap_or(0));
state
.store
.finalize_meeting(
&meeting_id,
FinalizeMeeting {
segments: segments.clone(),
speakers: meeting.speakers.clone(),
duration_secs,
recorded: meeting.recorded,
language: meeting.language,
backend_used: Some(backend.as_str().to_string()),
model_used: Some(model),
},
)
.await
.map_err(|e| WaError::new("storage", e.to_string()))?;
let notes_md = crate::notes::MarkdownNotes.to_markdown(&segments, &meeting.speakers, None);
let _ = std::fs::write(meeting_dir(&meeting_id).join("notes.md"), notes_md);
let _ = app.emit(
"transcript://finalized",
serde_json::json!({ "meetingId": meeting_id, "segmentCount": segments.len() }),
);
Ok(())
}
/// Re-run transcription from a `recovering` meeting's working `audio.wav`
@@ -346,12 +548,14 @@ pub async fn resume_transcription(
state: State<'_, AppState>,
meeting_id: MeetingId,
) -> WaResult<()> {
let model_path = whisper_model_path();
let settings = load_settings();
let model_id = model_id_for(&settings);
let model_path = whisper_model_file(&model_id);
if !model_path.exists() {
return Err(WaError::new(
"transcription",
format!(
"whisper model not found at {}; run `npm run download-models` first",
"whisper model '{model_id}' not found at {}; download it from Settings first",
model_path.display()
),
));
@@ -387,10 +591,6 @@ pub async fn resume_transcription(
.last()
.map(|s| (s.end_ms / 1000) as i64)
.unwrap_or(0);
let model_used = whisper_model_path()
.file_stem()
.and_then(|s| s.to_str())
.map(str::to_string);
state
.store
@@ -405,7 +605,7 @@ pub async fn resume_transcription(
recorded: true,
language: None,
backend_used: Some(BackendId::Cpu.as_str().to_string()),
model_used,
model_used: Some(model_id),
},
)
.await
+193 -5
View File
@@ -19,16 +19,204 @@ pub trait HardwareDetector: Send + Sync {
fn best(&self, preferred: Option<BackendId>) -> BackendInfo;
}
fn rank_of(id: BackendId) -> u8 {
match id {
BackendId::Npu => 0,
BackendId::Nvidia => 1,
BackendId::Amd => 2,
BackendId::Intel => 3,
BackendId::Cpu => 4,
}
}
fn cpu_backend() -> BackendInfo {
BackendInfo {
id: BackendId::Cpu,
name: "CPU".to_string(),
available: true,
rank: rank_of(BackendId::Cpu),
vram_mb: None,
}
}
/// Default Windows detector (DXGI for GPUs, ONNX/Windows ML for NPU).
pub struct WinHardwareDetector;
impl HardwareDetector for WinHardwareDetector {
fn detect(&self) -> Vec<BackendInfo> {
// T3.1: enumerate DXGI adapters (NVIDIA/AMD/Intel), probe NPU via ONNX/Windows ML,
// always include CPU. Return ranked list.
todo!("Phase 3 — detect backends")
let mut backends: Vec<BackendInfo> = Vec::new();
// ponytail: no NPU enumeration API is wired up yet — Windows ML/DirectML
// NPU discovery lands together with the real inference path (T3.4).
// Reporting it unavailable here (rather than guessing "present") keeps
// hardware_status honest until then.
backends.push(BackendInfo {
id: BackendId::Npu,
name: "NPU".to_string(),
available: false,
rank: rank_of(BackendId::Npu),
vram_mb: None,
});
#[cfg(windows)]
backends.extend(dxgi::enumerate_gpus());
for vendor in [BackendId::Nvidia, BackendId::Amd, BackendId::Intel] {
if !backends.iter().any(|b| b.id == vendor) {
backends.push(BackendInfo {
id: vendor,
name: format!("{vendor:?}"),
available: false,
rank: rank_of(vendor),
vram_mb: None,
});
}
}
backends.push(cpu_backend());
backends.sort_by_key(|b| b.rank);
backends
}
fn best(&self, _preferred: Option<BackendId>) -> BackendInfo {
todo!("Phase 3 — choose best backend / honor override")
fn best(&self, preferred: Option<BackendId>) -> BackendInfo {
pick_best(&self.detect(), preferred)
}
}
/// Pure selection logic, factored out of `best()` so it's testable without
/// real DXGI/hardware (an unavailable `preferred` falls through to the
/// highest-ranked available backend; CPU is the guaranteed last resort).
fn pick_best(backends: &[BackendInfo], preferred: Option<BackendId>) -> BackendInfo {
if let Some(pref) = preferred {
if let Some(b) = backends.iter().find(|b| b.id == pref && b.available) {
return b.clone();
}
}
backends
.iter()
.find(|b| b.available)
.cloned()
.unwrap_or_else(cpu_backend)
}
#[cfg(test)]
mod tests {
use super::*;
fn backend(id: BackendId, available: bool) -> BackendInfo {
BackendInfo {
id,
name: format!("{id:?}"),
available,
rank: rank_of(id),
vram_mb: None,
}
}
#[test]
fn falls_back_to_highest_ranked_available() {
let backends = vec![
backend(BackendId::Npu, false),
backend(BackendId::Nvidia, false),
backend(BackendId::Amd, true),
backend(BackendId::Cpu, true),
];
assert_eq!(pick_best(&backends, None).id, BackendId::Amd);
}
#[test]
fn ignores_unavailable_preference() {
let backends = vec![
backend(BackendId::Npu, false),
backend(BackendId::Cpu, true),
];
assert_eq!(
pick_best(&backends, Some(BackendId::Npu)).id,
BackendId::Cpu
);
}
#[test]
fn honors_available_preference_over_higher_rank() {
let backends = vec![
backend(BackendId::Nvidia, true),
backend(BackendId::Cpu, true),
];
assert_eq!(
pick_best(&backends, Some(BackendId::Cpu)).id,
BackendId::Cpu
);
}
#[test]
fn cpu_is_the_last_resort_even_if_absent_from_the_list() {
let backends = vec![backend(BackendId::Npu, false)];
assert_eq!(pick_best(&backends, None).id, BackendId::Cpu);
}
}
#[cfg(windows)]
mod dxgi {
use super::rank_of;
use crate::models::{BackendId, BackendInfo};
use windows::Win32::Graphics::Dxgi::{
CreateDXGIFactory1, IDXGIFactory1, DXGI_ADAPTER_FLAG_SOFTWARE,
};
fn backend_for_vendor(vendor_id: u32) -> Option<BackendId> {
match vendor_id {
0x10DE => Some(BackendId::Nvidia),
0x1002 | 0x1022 => Some(BackendId::Amd),
0x8086 => Some(BackendId::Intel),
_ => None,
}
}
/// Enumerates DXGI adapters, skipping software rasterizers (e.g. the
/// Microsoft Basic Render Driver), and maps each to a `BackendId` by PCI
/// vendor ID. Only the first adapter per vendor is kept — good enough to
/// answer "is an NVIDIA/AMD/Intel GPU present", which is all backend
/// selection needs; a multi-GPU picker is out of scope here.
pub fn enumerate_gpus() -> Vec<BackendInfo> {
let mut out = Vec::new();
let factory: windows::core::Result<IDXGIFactory1> = unsafe { CreateDXGIFactory1() };
let Ok(factory) = factory else {
tracing::warn!("DXGI factory creation failed; GPU backends unavailable");
return out;
};
let mut i = 0u32;
loop {
let adapter = unsafe { factory.EnumAdapters1(i) };
let adapter = match adapter {
Ok(a) => a,
Err(_) => break, // DXGI_ERROR_NOT_FOUND — enumeration exhausted
};
i += 1;
let Ok(desc) = (unsafe { adapter.GetDesc1() }) else {
continue;
};
if (desc.Flags & DXGI_ADAPTER_FLAG_SOFTWARE.0 as u32) != 0 {
continue;
}
let Some(id) = backend_for_vendor(desc.VendorId) else {
continue;
};
if out.iter().any(|b: &BackendInfo| b.id == id) {
continue;
}
let name = String::from_utf16_lossy(&desc.Description)
.trim_end_matches('\0')
.to_string();
out.push(BackendInfo {
id,
name,
available: true,
rank: rank_of(id),
vram_mb: Some((desc.DedicatedVideoMemory / (1024 * 1024)) as u32),
});
}
out
}
}
+10
View File
@@ -51,6 +51,11 @@ pub struct RecordingSession {
/// `transcript://segment` — lets `stop_recording` persist the full
/// transcript without a second round-trip through the event channel.
pub segments: Arc<StdMutex<Vec<models::TranscriptSegment>>>,
/// Backend actually in use, updated in place if the worker falls back
/// (T3.5); model id used for this session (T3.9's "low overhead" preset
/// can pick a different model than `Settings.whisper_model`).
pub active_backend: Arc<StdMutex<models::BackendId>>,
pub model_id: String,
}
/// Wraps the tray icon so it can be looked up from commands to update its
@@ -113,6 +118,11 @@ pub fn run() {
commands::acknowledge_recording_consent,
commands::resume_transcription,
commands::hardware_status,
commands::set_preferred_backend,
commands::list_models,
commands::download_model,
commands::remove_model,
commands::reprocess_transcript,
commands::list_meetings,
commands::get_meeting,
commands::delete_meeting,
+11
View File
@@ -37,6 +37,16 @@ pub struct BackendInfo {
pub vram_mb: Option<u32>,
}
/// A downloadable/installed Whisper model (T3.7, FR-MODEL-1).
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelInfo {
pub id: String, // e.g. "base.en-q5_1" — also the ggml filename stem
pub label: String,
pub size_mb: u32, // approximate download size
pub installed: bool,
pub active: bool,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum MeetingStatus {
@@ -143,6 +153,7 @@ pub struct Settings {
pub llm_endpoint: String,
pub llm_model: String,
pub preferred_backend: String, // auto|npu|nvidia|amd|intel|cpu
pub whisper_model: String, // ModelInfo.id, e.g. "base.en-q5_1"
pub low_overhead: bool,
// Recording retention (ADR-0009). Default OFF.
pub default_record: bool,
+13 -2
View File
@@ -29,6 +29,17 @@ pub fn meeting_dir(id: &MeetingId) -> PathBuf {
meetings_dir().join(id)
}
pub fn whisper_model_path() -> PathBuf {
wa_root().join("models").join("ggml-base.en-q5_1.bin")
pub fn models_dir() -> PathBuf {
wa_root().join("models")
}
/// Default model id used until the user picks one in Settings (T3.7).
pub const DEFAULT_WHISPER_MODEL: &str = "base.en-q5_1";
pub fn whisper_model_file(id: &str) -> PathBuf {
models_dir().join(format!("ggml-{id}.bin"))
}
pub fn whisper_model_path() -> PathBuf {
whisper_model_file(DEFAULT_WHISPER_MODEL)
}
+41 -7
View File
@@ -8,6 +8,8 @@ use crate::models::{BackendId, TranscriptSegment};
use std::path::Path;
use std::sync::mpsc::{Receiver, Sender};
pub mod models;
#[derive(Debug, thiserror::Error)]
pub enum TrxError {
#[error("model load failed: {0}")]
@@ -127,12 +129,18 @@ impl WhisperTranscriber {
#[cfg(feature = "cpu-transcription")]
impl Transcriber for WhisperTranscriber {
fn load(model: &Path, _backend: BackendId) -> Result<Self, TrxError> {
let ctx = whisper_rs::WhisperContext::new_with_params(
model,
whisper_rs::WhisperContextParameters::default(),
)
.map_err(|e| TrxError::Load(e.to_string()))?;
/// `backend` picks GPU offload at runtime (T3.2): whisper.cpp only has GPU
/// support at all when built with the `cuda`/`vulkan` Cargo feature (T3.3),
/// so `use_gpu` on a CPU-only build is a harmless no-op — this stays a
/// single code path either way rather than branching on which features
/// were compiled in.
fn load(model: &Path, backend: BackendId) -> Result<Self, TrxError> {
let params = whisper_rs::WhisperContextParameters {
use_gpu: !matches!(backend, BackendId::Cpu),
..Default::default()
};
let ctx = whisper_rs::WhisperContext::new_with_params(model, params)
.map_err(|e| TrxError::Load(e.to_string()))?;
Ok(Self {
ctx,
next_id: std::sync::atomic::AtomicU64::new(0),
@@ -168,10 +176,36 @@ fn available_threads() -> std::ffi::c_int {
.unwrap_or(4)
}
/// ONNX Runtime + DirectML transcriber for the NPU tier (Phase 3).
/// ONNX Runtime + DirectML transcriber for the NPU tier (T3.4).
///
/// ponytail: the encoder/decoder inference loop (ORT session + DirectML
/// execution provider, greedy decode over a Whisper ONNX export) is a
/// multi-day integration on its own and untestable without NPU hardware in
/// this environment — ships as a real `Transcriber` impl that fails loudly
/// instead of a silent no-op. `hardware::WinHardwareDetector` reports NPU as
/// unavailable (see hardware/mod.rs) so nothing routes here until it's built.
#[cfg(feature = "directml")]
pub struct OnnxNpuTranscriber;
#[cfg(feature = "directml")]
impl Transcriber for OnnxNpuTranscriber {
fn load(_model: &Path, _backend: BackendId) -> Result<Self, TrxError> {
Err(TrxError::Load(
"NPU (ONNX Runtime + DirectML) transcription is not yet implemented (T3.4)".into(),
))
}
fn transcribe_stream(&self, _audio: AudioWindow, _out: SegmentSink) -> Result<(), TrxError> {
Err(TrxError::Inference(
"NPU backend not yet implemented".into(),
))
}
fn transcribe_file(&self, _wav: &Path) -> Result<Vec<TranscriptSegment>, TrxError> {
Err(TrxError::Inference(
"NPU backend not yet implemented".into(),
))
}
}
/// Streaming window worker (Phase 1, T1.5/T1.6): accumulates raw 16kHz-mono
/// chunks from the `audio` service into fixed-size, **non-overlapping** windows
/// and runs one `transcribe_stream` pass per window as it fills, forwarding
+150
View File
@@ -0,0 +1,150 @@
//! Whisper model catalog + install/remove (Phase 3, T3.7, FR-MODEL-1).
//!
//! ponytail: a fixed list of known ggml quantized models, not a fetched index —
//! whisper.cpp's model set changes rarely enough that hardcoding it is the
//! lazy-correct choice; revisit only if we need custom/fine-tuned models.
use crate::models::ModelInfo;
use crate::paths::{models_dir, whisper_model_file};
use futures_util::StreamExt;
struct Catalog {
id: &'static str,
label: &'static str,
size_mb: u32,
}
const CATALOG: &[Catalog] = &[
Catalog {
id: "tiny.en-q5_1",
label: "Tiny (English, quantized) — fastest, least accurate",
size_mb: 32,
},
Catalog {
id: "base.en-q5_1",
label: "Base (English, quantized) — balanced default",
size_mb: 60,
},
Catalog {
id: "small.en-q5_1",
label: "Small (English, quantized) — more accurate, slower",
size_mb: 190,
},
Catalog {
id: "medium.en-q5_1",
label: "Medium (English, quantized) — best accuracy, slowest",
size_mb: 540,
},
];
fn model_url(id: &str) -> String {
format!("https://huggingface.co/ggerganov/whisper.cpp/resolve/main/ggml-{id}.bin")
}
pub fn list(active_id: &str) -> Vec<ModelInfo> {
CATALOG
.iter()
.map(|m| ModelInfo {
id: m.id.to_string(),
label: m.label.to_string(),
size_mb: m.size_mb,
installed: whisper_model_file(m.id).exists(),
active: m.id == active_id,
})
.collect()
}
#[derive(Debug, thiserror::Error)]
pub enum ModelError {
#[error("unknown model id: {0}")]
UnknownId(String),
#[error("download failed: {0}")]
Download(String),
#[error("{0}")]
Invalid(String),
#[error("io error: {0}")]
Io(#[from] std::io::Error),
}
/// Downloads a model, reporting progress via `on_progress(received, total)`.
/// Writes to a `.part` file first so a crash/cancel never leaves a truncated
/// model that `whisper_model_file` would treat as installed.
pub async fn download(
id: &str,
mut on_progress: impl FnMut(u64, Option<u64>),
) -> Result<(), ModelError> {
if !CATALOG.iter().any(|m| m.id == id) {
return Err(ModelError::UnknownId(id.to_string()));
}
std::fs::create_dir_all(models_dir())?;
let dest = whisper_model_file(id);
let tmp = dest.with_extension("part");
let resp = reqwest::get(model_url(id))
.await
.map_err(|e| ModelError::Download(e.to_string()))?;
if !resp.status().is_success() {
return Err(ModelError::Download(format!("HTTP {}", resp.status())));
}
let total = resp.content_length();
let mut received: u64 = 0;
let mut file = std::fs::File::create(&tmp)?;
let mut stream = resp.bytes_stream();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(|e| ModelError::Download(e.to_string()))?;
std::io::Write::write_all(&mut file, &chunk)?;
received += chunk.len() as u64;
on_progress(received, total);
}
drop(file);
std::fs::rename(&tmp, &dest)?;
Ok(())
}
pub fn remove(id: &str, active_id: &str) -> Result<(), ModelError> {
if id == active_id {
return Err(ModelError::Invalid(
"cannot remove the active model".to_string(),
));
}
let path = whisper_model_file(id);
if path.exists() {
std::fs::remove_file(path)?;
}
Ok(())
}
pub fn is_installed(id: &str) -> bool {
whisper_model_file(id).exists()
}
/// The smallest catalog model — used by the "low overhead" preset (T3.9).
pub fn smallest_id() -> &'static str {
CATALOG.first().map(|m| m.id).unwrap_or("base.en-q5_1")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn list_marks_exactly_the_active_id() {
let active = "small.en-q5_1";
let models = list(active);
assert_eq!(models.len(), CATALOG.len());
assert_eq!(models.iter().filter(|m| m.active).count(), 1);
assert!(models.iter().find(|m| m.id == active).unwrap().active);
}
#[test]
fn smallest_id_is_in_the_catalog() {
assert!(CATALOG.iter().any(|m| m.id == smallest_id()));
}
#[test]
fn remove_refuses_the_active_model_without_touching_disk() {
let active = smallest_id();
let err = remove(active, active).unwrap_err();
assert!(matches!(err, ModelError::Invalid(_)));
}
}