diff --git a/Cargo.lock b/Cargo.lock index 15d784a..e345558 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1624,6 +1624,17 @@ version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" +[[package]] +name = "socks" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0c3dbbd9ae980613c6dd8e28a9407b50509d3803b57624d5dfe8315218cd58b" +dependencies = [ + "byteorder", + "libc", + "winapi", +] + [[package]] name = "sqlite-wasm-rs" version = "0.5.5" @@ -2019,6 +2030,7 @@ dependencies = [ "once_cell", "rustls", "rustls-pki-types", + "socks", "url", "webpki-roots 0.26.11", ] @@ -2210,6 +2222,22 @@ dependencies = [ "rustls-pki-types", ] +[[package]] +name = "winapi" +version = "0.3.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5c839a674fcd7a98952e593242ea400abe93992746761e38641405d28b00f419" +dependencies = [ + "winapi-i686-pc-windows-gnu", + "winapi-x86_64-pc-windows-gnu", +] + +[[package]] +name = "winapi-i686-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac3b87c63620426dd9b991e5ce0329eff545bccbbb34f3be09ff6fb6ab51b7b6" + [[package]] name = "winapi-util" version = "0.1.11" @@ -2219,6 +2247,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "winapi-x86_64-pc-windows-gnu" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" + [[package]] name = "windows" version = "0.54.0" diff --git a/README.md b/README.md index 7ec0eba..8af6937 100644 --- a/README.md +++ b/README.md @@ -121,6 +121,7 @@ No mandatory engine downloads. No forced multi-gigabyte "summary engine" gating - **Pick any model** from the catalog — speed-first or accuracy-first — and switch per mode (Live vs File). - **Import your own speech models** — select a sherpa-onnx `.onnx` graph or Whisper GGML/GGUF file; Wisp validates it and automatically brings along matching tokens, encoder/decoder/joiner graphs, and external weights. +- **Downloads that work across regions** — automatic Hugging Face/mirror failover, resumable partial files, bounded retries, and environment, HTTP, or SOCKS proxy support. Change the mirror or proxy live under **Settings → Model downloads**; the same policy covers transcription, speaker, denoise, Core ML, and embedding models. - **Delete models** you don't need, right from the picker, to reclaim disk in one click. - **Configure everything** — transcription language, accuracy/speed profile, VAD gating, denoising, decoding thresholds, diarization, custom vocabulary — instead of being locked to one preset. - **Honest picker** — models your machine can't run are clearly marked, with size and hardware hints *before* you download. diff --git a/README.zh-Hans.md b/README.zh-Hans.md index bbc8a9e..ee6f7cd 100644 --- a/README.zh-Hans.md +++ b/README.zh-Hans.md @@ -121,6 +121,7 @@ Wisp 不只是「能在 Mac 上跑」—— 它贴着架构做了优化: - **任选模型** —— 从目录里挑,速度优先或精度优先 —— 还能按模式(实时 vs 文件)分别切换。 - **导入自己的语音模型** —— 选择 sherpa-onnx `.onnx` 图或 Whisper GGML/GGUF 文件;Wisp 会先验证格式,并自动带上匹配的 tokens、encoder/decoder/joiner 图和外部权重。 +- **跨地区可靠下载** —— 自动在 Hugging Face 与镜像之间切换,保留 partial 文件并断点续传,提供有限重试,以及环境变量、HTTP、SOCKS 代理支持。可在**设置 → 模型下载**中即时更换镜像或代理;转写、说话人、降噪、Core ML 和嵌入模型共用同一套策略。 - **删除模型** —— 不需要的直接在选择器里一键删掉,腾回磁盘空间。 - **一切可配** —— 转写语言、精度/速度档位、VAD 门控、降噪、解码阈值、说话人分离、自定义词表 —— 而不是被锁死在某个预设上。 - **诚实的选择器** —— 你的机器跑不动的模型会被清楚标出,下载*之前*就给出体积和硬件提示。 diff --git a/app/src-tauri/Cargo.lock b/app/src-tauri/Cargo.lock index e6e8dc7..6ed6c5d 100644 --- a/app/src-tauri/Cargo.lock +++ b/app/src-tauri/Cargo.lock @@ -6750,6 +6750,7 @@ dependencies = [ "tokenizers", "ureq", "wisp-library", + "wisp-models", ] [[package]] diff --git a/app/src-tauri/src/lib.rs b/app/src-tauri/src/lib.rs index 3d679e9..17440f1 100644 --- a/app/src-tauri/src/lib.rs +++ b/app/src-tauri/src/lib.rs @@ -50,7 +50,8 @@ use wisp_loopback::WasapiLoopbackSource; use wisp_models::{ builtin_catalog, cloud_catalog, coreml_asset, denoise_models, diarization_models, family_runnable, model_fit, recommended_accurate_model, recommended_default_model, Accelerator, - FsModelStore, GpuTier, HttpDownloader, MachineProfile, ModelFit, + DownloadConfig, DownloadSource, FsModelStore, GpuTier, HttpDownloader, MachineProfile, + ModelFit, ProxyMode, }; use wisp_pipeline::{ remap_to_original, transcribe_in_windows, EnergySegmenter, EnergyVad, GatedClip, LiveStream, @@ -113,6 +114,10 @@ const MIC_VAD_THRESHOLD: f32 = 0.012; /// Shared application state. struct AppState { store: Arc, + /// Shared by every local-model path (ASR/support/Core ML/embeddings), and live-updated from + /// Settings without rebuilding the model store. + downloader: HttpDownloader, + download_settings_path: PathBuf, sessions: Mutex>, /// The push-to-talk dictation session, while the hotkey is held; `None` otherwise. dictation: Mutex>, @@ -237,6 +242,104 @@ struct DownloadProgressDto { total: u64, } +/// User-facing regional download settings persisted in app data. Runtime-only retry/timeouts stay +/// in [`DownloadConfig`] so the public settings surface remains small and stable. +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +struct DownloadSettingsDto { + source: String, + mirror_url: String, + proxy_mode: String, + proxy_url: String, +} + +impl Default for DownloadSettingsDto { + fn default() -> Self { + Self::from_config(&DownloadConfig::default()) + } +} + +impl DownloadSettingsDto { + fn from_config(config: &DownloadConfig) -> Self { + Self { + source: match config.source { + DownloadSource::Auto => "auto", + DownloadSource::Official => "official", + DownloadSource::MirrorFirst => "mirror", + } + .to_owned(), + mirror_url: config.mirror_url.clone(), + proxy_mode: match config.proxy_mode { + ProxyMode::System => "system", + ProxyMode::Direct => "direct", + ProxyMode::Custom => "custom", + } + .to_owned(), + proxy_url: config.proxy_url.clone(), + } + } + + fn to_config(&self) -> Result { + let source = match self.source.as_str() { + "auto" => DownloadSource::Auto, + "official" => DownloadSource::Official, + "mirror" => DownloadSource::MirrorFirst, + other => return Err(format!("unknown download source: {other}")), + }; + let proxy_mode = match self.proxy_mode.as_str() { + "system" => ProxyMode::System, + "direct" => ProxyMode::Direct, + "custom" => ProxyMode::Custom, + other => return Err(format!("unknown proxy mode: {other}")), + }; + let config = DownloadConfig { + source, + mirror_url: self.mirror_url.trim().trim_end_matches('/').to_owned(), + proxy_mode, + proxy_url: self.proxy_url.trim().to_owned(), + ..DownloadConfig::default() + }; + config.validate().map_err(|error| error.to_string())?; + Ok(config) + } +} + +fn load_download_settings(path: &Path) -> DownloadSettingsDto { + fs::read_to_string(path) + .ok() + .and_then(|json| serde_json::from_str::(&json).ok()) + .filter(|settings| settings.to_config().is_ok()) + .unwrap_or_default() +} + +fn save_download_settings(path: &Path, settings: &DownloadSettingsDto) -> Result<(), String> { + let json = serde_json::to_string_pretty(settings).map_err(|error| error.to_string())?; + fs::write(path, json).map_err(|error| format!("save download settings: {error}")) +} + +#[tauri::command] +fn get_download_settings(state: State<'_, AppState>) -> Result { + let config = state + .downloader + .config() + .map_err(|error| error.to_string())?; + Ok(DownloadSettingsDto::from_config(&config)) +} + +#[tauri::command] +fn set_download_settings( + state: State<'_, AppState>, + settings: DownloadSettingsDto, +) -> Result { + let config = settings.to_config()?; + save_download_settings(&state.download_settings_path, &settings)?; + state + .downloader + .update_config(config) + .map_err(|error| error.to_string())?; + Ok(settings) +} + /// Metadata emitted at the start of a file transcription (for the progress bar). #[derive(Clone, Serialize)] #[serde(rename_all = "camelCase")] @@ -4700,7 +4803,9 @@ fn spawn_embed_progress_poller( /// frontend activates via [`set_embedding_model`] afterward, mirroring the ASR picker. #[tauri::command] async fn download_embedding_model(app: AppHandle, id: String) -> Result<(), String> { - let cache_root = app.state::().embed_cache_dir.clone(); + let state = app.state::(); + let cache_root = state.embed_cache_dir.clone(); + let downloader = state.downloader.clone(); tauri::async_runtime::spawn_blocking(move || { let model = wisp_embed::catalog_model(&id) @@ -4719,7 +4824,7 @@ async fn download_embedding_model(app: AppHandle, id: String) -> Result<(), Stri Arc::clone(&stop), ); - let downloaded = wisp_embed::download_model(model, &dir); + let downloaded = wisp_embed::download_model_with(model, &dir, &downloader); stop.store(true, Ordering::Relaxed); let _ = poller.join(); @@ -4874,6 +4979,12 @@ pub fn run() { let cloud_custom_models = load_cloud_custom_models(&cloud_custom_models_path); let cloud_custom_endpoints_path = data_dir.join("cloud-custom-endpoints.json"); let cloud_custom_endpoints = load_cloud_endpoints(&cloud_custom_endpoints_path); + let download_settings_path = data_dir.join("download-settings.json"); + let download_config = load_download_settings(&download_settings_path) + .to_config() + .unwrap_or_default(); + let downloader = HttpDownloader::with_config(download_config) + .map_err(|error| std::io::Error::other(error.to_string()))?; let custom_models_dir = data_dir.join("custom-models"); let _ = fs::create_dir_all(&custom_models_dir); let custom_models = load_custom_models(&custom_models_dir); @@ -4883,7 +4994,7 @@ pub fn run() { let store = Arc::new(FsModelStore::new( data_dir.join("models"), [asr_catalog.clone(), diarization_models(), denoise_models()].concat(), - Box::new(HttpDownloader), + Box::new(downloader.clone()), )); // Prefer the last chosen model (if still installed), then the first installed, then the @@ -4931,6 +5042,8 @@ pub fn run() { app.manage(AppState { store, + downloader, + download_settings_path, sessions: Mutex::new(Vec::new()), dictation: Mutex::new(None), dictation_hotkey: Mutex::new(dictation::DEFAULT_DICTATION_HOTKEY.to_owned()), @@ -5003,6 +5116,8 @@ pub fn run() { assist_realtime_params, run_llm_task, run_assist_stream, + get_download_settings, + set_download_settings, download_model, import_custom_model, download_coreml, @@ -5685,4 +5800,39 @@ mod tests { let _ = fs::remove_dir_all(&dir); } + + #[test] + fn download_settings_round_trip_and_map_to_runtime_policy() { + let path = temp_path("download-settings"); + let settings = DownloadSettingsDto { + source: "mirror".to_owned(), + mirror_url: "https://models.example.cn/hf".to_owned(), + proxy_mode: "custom".to_owned(), + proxy_url: "socks5://127.0.0.1:1080".to_owned(), + }; + + save_download_settings(&path, &settings).unwrap(); + assert_eq!(load_download_settings(&path), settings); + let config = settings.to_config().unwrap(); + assert_eq!(config.source, DownloadSource::MirrorFirst); + assert_eq!(config.proxy_mode, ProxyMode::Custom); + assert_eq!(config.mirror_url, "https://models.example.cn/hf"); + + let _ = fs::remove_file(path); + } + + #[test] + fn invalid_download_settings_are_rejected_without_persisting() { + let unknown_source = DownloadSettingsDto { + source: "nearest-magic".to_owned(), + ..DownloadSettingsDto::default() + }; + assert!(unknown_source.to_config().is_err()); + + let insecure_mirror = DownloadSettingsDto { + mirror_url: "http://mirror.invalid".to_owned(), + ..DownloadSettingsDto::default() + }; + assert!(insecure_mirror.to_config().is_err()); + } } diff --git a/app/src/lib/Settings.svelte b/app/src/lib/Settings.svelte index 4ec23ea..365e010 100644 --- a/app/src/lib/Settings.svelte +++ b/app/src/lib/Settings.svelte @@ -13,10 +13,11 @@ let { open = $bindable(false), autoSave = $bindable(false) }: { open?: boolean; autoSave?: boolean } = $props(); - type Section = "models" | "search" | "dictation" | "storage"; + type Section = "models" | "search" | "downloads" | "dictation" | "storage"; const sections = $derived<{ id: Section; label: string }[]>([ { id: "models", label: i18n.t.settings.aiModels }, { id: "search", label: i18n.t.settings.search }, + { id: "downloads", label: i18n.t.settings.downloads }, { id: "dictation", label: i18n.t.settings.dictation }, { id: "storage", label: i18n.t.settings.storage }, ]); @@ -70,6 +71,45 @@ } } + // ── Regional downloads ───────────────────────────────────────────────────────────────────────── + type DownloadSettings = { + source: "auto" | "official" | "mirror"; + mirrorUrl: string; + proxyMode: "system" | "direct" | "custom"; + proxyUrl: string; + }; + let downloads = $state({ + source: "auto", + mirrorUrl: "https://hf-mirror.com", + proxyMode: "system", + proxyUrl: "", + }); + let downloadSettingsBusy = $state(false); + let downloadSettingsError = $state(""); + let downloadSettingsSaved = $state(false); + + async function loadDownloadSettings() { + try { + downloads = await invoke("get_download_settings"); + downloadSettingsError = ""; + } catch (e) { + downloadSettingsError = String(e); + } + } + + async function saveDownloadSettings() { + downloadSettingsBusy = true; + downloadSettingsSaved = false; + downloadSettingsError = ""; + try { + downloads = await invoke("set_download_settings", { settings: downloads }); + downloadSettingsSaved = true; + } catch (e) { + downloadSettingsError = String(e); + } + downloadSettingsBusy = false; + } + // ── Storage category ──────────────────────────────────────────────────────────────────────────── // Where Wisp keeps things on disk, resolved from the Tauri app-data dir — the same dir the Rust side // opens the SQLite library and the model store under. @@ -101,6 +141,7 @@ $effect(() => { if (open) { loadDictation(); + loadDownloadSettings(); loadPaths(); } }); @@ -154,6 +195,85 @@ {:else if section === "search"} + {:else if section === "downloads"} +

{i18n.t.settings.downloadsIntro}

+ + + + + + + + {#if downloads.proxyMode === "custom"} + + {/if} + +
+ + {#if downloadSettingsSaved}{i18n.t.settings.downloadSettingsSaved}{/if} +
+ {#if downloadSettingsError}

{downloadSettingsError}

{/if} {:else if section === "dictation"}

{i18n.t.settings.dictationIntro}

@@ -389,6 +509,60 @@ border-color: var(--accent); } + .set-btn.primary { + color: var(--accent); + border-color: var(--accent); + } + + .set-btn:disabled { + opacity: 0.55; + cursor: default; + } + + .set-field { + display: flex; + flex-direction: column; + gap: 6px; + } + + .set-input { + width: 100%; + box-sizing: border-box; + font-family: var(--font-mono); + font-size: 12px; + color: var(--text); + background: var(--surface); + border: 1px solid var(--border-strong); + border-radius: 7px; + padding: 7px 9px; + } + + .set-input:disabled { + opacity: 0.5; + } + + .set-select { + width: 230px; + font-family: inherit; + } + + .set-help { + font-size: 11.5px; + line-height: 1.45; + color: var(--muted); + } + + .set-actions { + display: flex; + align-items: center; + gap: 10px; + } + + .set-saved { + font-size: 12px; + color: var(--accent); + } + .hotkey-input { font-family: var(--font-mono); font-size: 12px; diff --git a/app/src/lib/i18n/en.ts b/app/src/lib/i18n/en.ts index 848512e..a1c4323 100644 --- a/app/src/lib/i18n/en.ts +++ b/app/src/lib/i18n/en.ts @@ -275,6 +275,25 @@ export const en = { settings: { aiModels: "AI models", search: "Notes search", + downloads: "Model downloads", + downloadsIntro: + "Reliable downloads for every region. Automatic mode fails over between Hugging Face and your mirror; interrupted files resume instead of restarting.", + downloadSource: "Source", + downloadSourceAuto: "Automatic · fail over", + downloadSourceOfficial: "Hugging Face only", + downloadSourceMirror: "Mirror first", + downloadMirror: "Hugging Face mirror", + downloadMirrorHint: + "Used for Hugging Face files only. The default is a third-party public mirror for mainland China; use your organisation's HTTPS mirror when required.", + downloadProxy: "Proxy", + downloadProxySystem: "Environment variables", + downloadProxyDirect: "Direct connection", + downloadProxyCustom: "Custom proxy", + downloadProxyUrl: "Proxy URL", + downloadProxyHint: "Supports HTTP, SOCKS4, SOCKS4A, and SOCKS5, including authenticated URLs.", + downloadSettingsSave: "Save download settings", + downloadSettingsSaving: "Saving…", + downloadSettingsSaved: "Saved · applies to the next request", searchIntro: "How the Library searches your saved notes. Semantic and Hybrid understand meaning (not just exact words) — they need a local embedding model, downloaded once and run fully on-device.", searchMode: "Search mode", diff --git a/app/src/lib/i18n/zh-Hans.ts b/app/src/lib/i18n/zh-Hans.ts index 1c2a618..d81ce4c 100644 --- a/app/src/lib/i18n/zh-Hans.ts +++ b/app/src/lib/i18n/zh-Hans.ts @@ -242,6 +242,25 @@ export const zhHans: Messages = { settings: { aiModels: "AI 模型", search: "笔记搜索", + downloads: "模型下载", + downloadsIntro: + "为不同地区提供更可靠的下载。自动模式会在 Hugging Face 与镜像之间切换;网络中断后会续传,不再从零开始。", + downloadSource: "下载源", + downloadSourceAuto: "自动 · 失败时切换", + downloadSourceOfficial: "仅 Hugging Face", + downloadSourceMirror: "镜像优先", + downloadMirror: "Hugging Face 镜像", + downloadMirrorHint: + "只用于 Hugging Face 模型文件。默认值是面向中国大陆的第三方公共镜像;有合规要求时请填写组织自己的 HTTPS 镜像。", + downloadProxy: "代理", + downloadProxySystem: "环境变量", + downloadProxyDirect: "直接连接", + downloadProxyCustom: "自定义代理", + downloadProxyUrl: "代理 URL", + downloadProxyHint: "支持 HTTP、SOCKS4、SOCKS4A 和 SOCKS5,也支持带认证信息的 URL。", + downloadSettingsSave: "保存下载设置", + downloadSettingsSaving: "保存中…", + downloadSettingsSaved: "已保存 · 下次请求立即生效", searchIntro: "笔记库如何搜索你保存的笔记。语义与混合搜索理解语义(不只是字面匹配)—— 需要一个本地嵌入模型,下载一次后全程在设备端运行。", searchMode: "搜索模式", diff --git a/app/src/lib/i18n/zh-Hant.ts b/app/src/lib/i18n/zh-Hant.ts index 8cc68e0..a28f2c3 100644 --- a/app/src/lib/i18n/zh-Hant.ts +++ b/app/src/lib/i18n/zh-Hant.ts @@ -242,6 +242,25 @@ export const zhHant: Messages = { settings: { aiModels: "AI 模型", search: "筆記搜尋", + downloads: "模型下載", + downloadsIntro: + "為不同地區提供更可靠的下載。自動模式會在 Hugging Face 與鏡像之間切換;網路中斷後會續傳,不再從零開始。", + downloadSource: "下載來源", + downloadSourceAuto: "自動 · 失敗時切換", + downloadSourceOfficial: "僅 Hugging Face", + downloadSourceMirror: "鏡像優先", + downloadMirror: "Hugging Face 鏡像", + downloadMirrorHint: + "只用於 Hugging Face 模型檔案。預設值是面向中國大陸的第三方公共鏡像;有合規要求時請填寫組織自己的 HTTPS 鏡像。", + downloadProxy: "代理", + downloadProxySystem: "環境變數", + downloadProxyDirect: "直接連線", + downloadProxyCustom: "自訂代理", + downloadProxyUrl: "代理 URL", + downloadProxyHint: "支援 HTTP、SOCKS4、SOCKS4A 和 SOCKS5,也支援帶驗證資訊的 URL。", + downloadSettingsSave: "儲存下載設定", + downloadSettingsSaving: "儲存中…", + downloadSettingsSaved: "已儲存 · 下次請求立即生效", searchIntro: "筆記庫如何搜尋你儲存的筆記。語意與混合搜尋理解語意(不只是字面比對)—— 需要一個本機嵌入模型,下載一次後全程在裝置端執行。", searchMode: "搜尋模式", diff --git a/crates/wisp-embed/Cargo.lock b/crates/wisp-embed/Cargo.lock index 4625ccb..4dda4a5 100644 --- a/crates/wisp-embed/Cargo.lock +++ b/crates/wisp-embed/Cargo.lock @@ -31,6 +31,15 @@ dependencies = [ "memchr", ] +[[package]] +name = "arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1" +dependencies = [ + "derive_arbitrary", +] + [[package]] name = "autocfg" version = "1.5.1" @@ -219,6 +228,17 @@ dependencies = [ "serde", ] +[[package]] +name = "derive_arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "derive_builder" version = "0.20.2" @@ -277,6 +297,12 @@ version = "1.16.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + [[package]] name = "errno" version = "0.3.14" @@ -305,6 +331,12 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7360491ce676a36bf9bb3c56c1aa791658183a54d2744120f27285738d90465a" +[[package]] +name = "fastrand" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" + [[package]] name = "filetime" version = "0.2.29" @@ -532,6 +564,16 @@ dependencies = [ "icu_properties", ] +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown 0.17.1", +] + [[package]] name = "itertools" version = "0.14.0" @@ -1207,6 +1249,19 @@ dependencies = [ "xattr", ] +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.3.4", + "once_cell", + "rustix", + "windows-sys 0.61.2", +] + [[package]] name = "thiserror" version = "2.0.18" @@ -1579,9 +1634,12 @@ dependencies = [ "ndarray", "ort", "serde_json", + "tempfile", "tokenizers", "ureq", + "wisp-core", "wisp-library", + "wisp-models", ] [[package]] @@ -1594,6 +1652,16 @@ dependencies = [ "wisp-core", ] +[[package]] +name = "wisp-models" +version = "0.0.0" +dependencies = [ + "sha2", + "ureq", + "wisp-core", + "zip", +] + [[package]] name = "wit-bindgen" version = "0.57.1" @@ -1719,8 +1787,37 @@ dependencies = [ "syn", ] +[[package]] +name = "zip" +version = "2.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50" +dependencies = [ + "arbitrary", + "crc32fast", + "crossbeam-utils", + "displaydoc", + "flate2", + "indexmap", + "memchr", + "thiserror", + "zopfli", +] + [[package]] name = "zmij" version = "1.0.21" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" + +[[package]] +name = "zopfli" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249" +dependencies = [ + "bumpalo", + "crc32fast", + "log", + "simd-adler32", +] diff --git a/crates/wisp-embed/Cargo.toml b/crates/wisp-embed/Cargo.toml index 3488d89..2235a8c 100644 --- a/crates/wisp-embed/Cargo.toml +++ b/crates/wisp-embed/Cargo.toml @@ -8,6 +8,7 @@ license = "MIT" [dependencies] # The Embedder contract this crate implements. wisp-library = { path = "../wisp-library" } +wisp-models = { path = "../wisp-models" } # Runtime-free HTTP client for self-downloading local model files from the HF hub and for the cloud # embedder's OpenAI-compatible /embeddings calls. `ureq` has no async runtime of its own, so it is # safe to call from a plain Tauri command thread (unlike reqwest::blocking). @@ -21,3 +22,7 @@ serde_json = "1" ort = { version = "=2.0.0-rc.9", features = ["ndarray", "download-binaries"] } tokenizers = { version = "0.21", default-features = false, features = ["onig"] } ndarray = "0.16" + +[dev-dependencies] +tempfile = "3" +wisp-core = { path = "../wisp-core" } diff --git a/crates/wisp-embed/src/lib.rs b/crates/wisp-embed/src/lib.rs index 869fd45..7b2ded8 100644 --- a/crates/wisp-embed/src/lib.rs +++ b/crates/wisp-embed/src/lib.rs @@ -9,6 +9,7 @@ use std::path::{Path, PathBuf}; use wisp_library::{LibraryError, Result}; +use wisp_models::{FileDownloader, HttpDownloader}; mod cloud; pub use cloud::{cloud_catalog_model, CloudCatalogModel, CloudEmbedder, CLOUD_CATALOG}; @@ -240,6 +241,16 @@ pub fn catalog_model(id: &str) -> Option<&'static CatalogModel> { /// interrupted download never leaves a half-written file behind; already-present files are skipped, /// making a re-run resume. Callers observe progress by watching `dir` grow. pub fn download_model(model: &CatalogModel, dir: &Path) -> Result<()> { + download_model_with(model, dir, &HttpDownloader::default()) +} + +/// Downloads with Wisp's shared regional downloader, so the app can apply its live mirror/proxy +/// settings to embedding models as well as transcription models. +pub fn download_model_with( + model: &CatalogModel, + dir: &Path, + downloader: &dyn FileDownloader, +) -> Result<()> { for file in model.files { let dest = dir.join(file); if dest.exists() { @@ -251,47 +262,22 @@ pub fn download_model(model: &CatalogModel, dir: &Path) -> Result<()> { } let url = format!("https://huggingface.co/{}/resolve/main/{file}", model.repo); - fetch_atomic(&url, &dest)?; + fetch_atomic(&url, &dest, downloader)?; } Ok(()) } /// Streams `url` into `dest` via a sibling `*.part`, renamed into place only after the whole body is -/// written; a failed transfer removes the `*.part` so it never lingers. -fn fetch_atomic(url: &str, dest: &Path) -> Result<()> { +/// written. A failed transfer retains only the non-final `*.part`, allowing the next attempt to +/// continue with an HTTP Range request. +fn fetch_atomic(url: &str, dest: &Path, downloader: &dyn FileDownloader) -> Result<()> { let mut part = dest.as_os_str().to_owned(); part.push(".part"); let part = PathBuf::from(part); - let streamed = (|| -> Result<()> { - let resp = ureq::get(url) - .call() - .map_err(|e| LibraryError::Embed(format!("download {url}: {e}")))?; - - // Hold the transfer to the declared length. Without this, a connection dropped mid-body on a - // length-less (close-delimited) response reads as a clean EOF, and the truncated file gets - // renamed into place as "complete" — then it fails to load forever, since the present (short) - // file makes the download skip it on every retry. - let expected: Option = resp.header("Content-Length").and_then(|h| h.parse().ok()); - - let mut reader = resp.into_reader(); - let mut out = std::fs::File::create(&part).map_err(io_err)?; - let written = std::io::copy(&mut reader, &mut out).map_err(io_err)?; - - if let Some(total) = expected { - if written != total { - return Err(LibraryError::Embed(format!( - "download {url}: truncated ({written} of {total} bytes)" - ))); - } - } - Ok(()) - })(); - - if let Err(e) = streamed { - let _ = std::fs::remove_file(&part); - return Err(e); - } + downloader + .download(url, &part) + .map_err(|error| LibraryError::Embed(error.to_string()))?; std::fs::rename(&part, dest).map_err(io_err) } @@ -303,6 +289,19 @@ fn io_err(e: std::io::Error) -> LibraryError { #[cfg(test)] mod tests { use super::*; + use std::sync::{Arc, Mutex}; + use wisp_core::error::Result as ModelResult; + use wisp_models::FileDownloader; + + struct RecordingDownloader(Arc>>); + + impl FileDownloader for RecordingDownloader { + fn download(&self, url: &str, dest: &Path) -> ModelResult<()> { + self.0.lock().unwrap().push(url.to_owned()); + std::fs::write(dest, url)?; + Ok(()) + } + } #[test] fn catalog_ids_are_unique_and_findable() { @@ -314,6 +313,20 @@ mod tests { assert!(catalog_model("nope").is_none()); } + #[test] + fn catalog_download_can_use_the_shared_regional_downloader() { + let temp = tempfile::tempdir().unwrap(); + let calls = Arc::new(Mutex::new(Vec::new())); + let downloader = RecordingDownloader(Arc::clone(&calls)); + + download_model_with(&CATALOG[0], temp.path(), &downloader).unwrap(); + + assert_eq!(calls.lock().unwrap().len(), CATALOG[0].files.len()); + for file in CATALOG[0].files { + assert!(temp.path().join(file).is_file(), "{file}"); + } + } + #[test] fn catalog_prefixes_match_family() { // E5 is asymmetric (both prefixes set); BGE-zh v1.5 is symmetric (no instruction). diff --git a/crates/wisp-models/Cargo.toml b/crates/wisp-models/Cargo.toml index 46d4c2f..f1dac65 100644 --- a/crates/wisp-models/Cargo.toml +++ b/crates/wisp-models/Cargo.toml @@ -19,7 +19,7 @@ whisper-vulkan = [] [dependencies] wisp-core = { path = "../wisp-core" } sha2 = "0.10" -ureq = { version = "2", optional = true } +ureq = { version = "2", optional = true, features = ["proxy-from-env", "socks-proxy"] } # Unzips the Core ML encoder archive (a `.mlmodelc` directory) next to a whisper.cpp model. zip = { version = "2", default-features = false, features = ["deflate"] } diff --git a/crates/wisp-models/src/download.rs b/crates/wisp-models/src/download.rs index db881cc..22d823b 100644 --- a/crates/wisp-models/src/download.rs +++ b/crates/wisp-models/src/download.rs @@ -1,8 +1,98 @@ //! Pluggable file downloading. +use std::fs::{File, OpenOptions}; +use std::io::{Read, Write}; use std::path::Path; +use std::sync::{Arc, RwLock}; +use std::time::Duration; -use wisp_core::error::Result; +use wisp_core::error::{Result, WispError}; + +/// Public Hugging Face endpoint commonly reachable from mainland China. Users can replace this +/// with an organisation-controlled mirror in Settings. +pub const DEFAULT_HF_MIRROR: &str = "https://hf-mirror.com"; + +/// How Hugging Face downloads choose between the official endpoint and the configured mirror. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum DownloadSource { + /// Try the endpoint that last succeeded first, failing over automatically. + #[default] + Auto, + /// Use the official endpoint only. + Official, + /// Try the configured mirror first, then fall back to the official endpoint. + MirrorFirst, +} + +/// Proxy selection for model downloads. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum ProxyMode { + /// Honour `ALL_PROXY`, `HTTPS_PROXY`, and `HTTP_PROXY` when inherited by the app. + #[default] + System, + /// Bypass proxy environment variables. + Direct, + /// Route downloads through the explicit HTTP/SOCKS proxy URL. + Custom, +} + +/// Runtime model-download policy. This is intentionally independent from app persistence/UI types, +/// so the downloader stays reusable in tests and other Wisp frontends. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct DownloadConfig { + pub source: DownloadSource, + pub mirror_url: String, + pub proxy_mode: ProxyMode, + pub proxy_url: String, + pub connect_timeout_secs: u64, + pub read_timeout_secs: u64, + pub retries: u8, +} + +impl Default for DownloadConfig { + fn default() -> Self { + Self { + source: DownloadSource::Auto, + mirror_url: DEFAULT_HF_MIRROR.to_owned(), + proxy_mode: ProxyMode::System, + proxy_url: String::new(), + connect_timeout_secs: 8, + read_timeout_secs: 45, + retries: 3, + } + } +} + +impl DownloadConfig { + pub fn validate(&self) -> Result<()> { + if !self.mirror_url.starts_with("https://") { + return Err(WispError::Model( + "model mirror must use an https:// URL".to_owned(), + )); + } + if self.proxy_mode == ProxyMode::Custom && self.proxy_url.trim().is_empty() { + return Err(WispError::Model( + "enter an HTTP or SOCKS proxy URL for custom proxy mode".to_owned(), + )); + } + if self.proxy_mode == ProxyMode::Custom { + ureq::Proxy::new(self.proxy_url.trim()) + .map_err(|error| WispError::Model(format!("invalid proxy URL: {error}")))?; + } + if self.connect_timeout_secs == 0 || self.read_timeout_secs == 0 || self.retries == 0 { + return Err(WispError::Model( + "download timeouts and retry count must be greater than zero".to_owned(), + )); + } + Ok(()) + } +} + +#[derive(Debug)] +struct DownloadState { + config: DownloadConfig, + preferred_mirror: Option, +} /// Downloads a single file from a URL to a local path. /// @@ -26,10 +116,111 @@ pub trait FileDownloader: Send + Sync { } } -/// A [`FileDownloader`] backed by `ureq` over HTTPS. +/// A regional, resumable [`FileDownloader`] backed by `ureq` over HTTPS. #[cfg(feature = "http")] -#[derive(Debug, Default)] -pub struct HttpDownloader; +#[derive(Clone, Debug)] +pub struct HttpDownloader { + state: Arc>, +} + +#[cfg(feature = "http")] +impl Default for HttpDownloader { + fn default() -> Self { + Self::with_config(DownloadConfig::default()) + .expect("the built-in download configuration is valid") + } +} + +#[cfg(feature = "http")] +impl HttpDownloader { + pub fn with_config(config: DownloadConfig) -> Result { + config.validate()?; + Ok(Self { + state: Arc::new(RwLock::new(DownloadState { + config, + preferred_mirror: None, + })), + }) + } + + pub fn config(&self) -> Result { + self.state + .read() + .map(|state| state.config.clone()) + .map_err(|_| WispError::Model("download settings lock poisoned".to_owned())) + } + + pub fn update_config(&self, config: DownloadConfig) -> Result<()> { + config.validate()?; + let mut state = self + .state + .write() + .map_err(|_| WispError::Model("download settings lock poisoned".to_owned()))?; + state.config = config; + state.preferred_mirror = None; + Ok(()) + } + + fn candidate_urls(&self, url: &str) -> Vec { + let Ok(state) = self.state.read() else { + return vec![url.to_owned()]; + }; + candidate_urls(url, &state.config, state.preferred_mirror) + } + + fn remember_success(&self, original: &str, successful: &str) { + if !original.starts_with("https://huggingface.co/") { + return; + } + if let Ok(mut state) = self.state.write() { + state.preferred_mirror = Some(successful != original); + } + } + + fn agent(config: &DownloadConfig) -> Result { + let mut builder = ureq::AgentBuilder::new() + .timeout_connect(Duration::from_secs(config.connect_timeout_secs)) + .timeout_read(Duration::from_secs(config.read_timeout_secs)) + .timeout_write(Duration::from_secs(config.read_timeout_secs)) + .redirects(10) + .user_agent(concat!("Wisp/", env!("CARGO_PKG_VERSION"))); + + builder = match config.proxy_mode { + ProxyMode::System => builder.try_proxy_from_env(true), + ProxyMode::Direct => builder.try_proxy_from_env(false), + ProxyMode::Custom => { + let proxy = ureq::Proxy::new(config.proxy_url.trim()) + .map_err(|error| WispError::Model(format!("invalid proxy URL: {error}")))?; + builder.proxy(proxy) + } + }; + Ok(builder.build()) + } +} + +fn candidate_urls( + url: &str, + config: &DownloadConfig, + preferred_mirror: Option, +) -> Vec { + let Some(path) = url.strip_prefix("https://huggingface.co/") else { + return vec![url.to_owned()]; + }; + if config.source == DownloadSource::Official { + return vec![url.to_owned()]; + } + + let mirror = format!("{}/{}", config.mirror_url.trim_end_matches('/'), path); + let mirror_first = + config.source == DownloadSource::MirrorFirst || preferred_mirror == Some(true); + let mut urls = if mirror_first { + vec![mirror, url.to_owned()] + } else { + vec![url.to_owned(), mirror] + }; + urls.dedup(); + urls +} #[cfg(feature = "http")] impl FileDownloader for HttpDownloader { @@ -43,28 +234,333 @@ impl FileDownloader for HttpDownloader { dest: &Path, on_bytes: &mut dyn FnMut(u64), ) -> Result<()> { - use std::fs::File; - use std::io::{Read, Write}; - use wisp_core::error::WispError; - - let response = ureq::get(url) - .call() - .map_err(|e| WispError::Model(format!("download {url}: {e}")))?; - - let mut reader = response.into_reader(); - let mut file = File::create(dest)?; - let mut buf = vec![0u8; 64 * 1024]; - let mut written = 0u64; - - loop { - let n = reader.read(&mut buf)?; - if n == 0 { - break; + let config = self.config()?; + let agent = Self::agent(&config)?; + let candidates = self.candidate_urls(url); + let mut errors = Vec::new(); + + for round in 0..config.retries { + for candidate in &candidates { + match transfer_once(&agent, candidate, dest, on_bytes) { + Ok(()) => { + self.remember_success(url, candidate); + return Ok(()); + } + Err(error) => errors.push(format!("{candidate}: {error}")), + } + } + if round + 1 < config.retries { + let delay = 300u64.saturating_mul(1u64 << round.min(4)); + std::thread::sleep(Duration::from_millis(delay)); } - file.write_all(&buf[..n])?; - written += n as u64; - on_bytes(written); } - Ok(()) + + let attempted = errors + .iter() + .rev() + .take(candidates.len().max(1)) + .cloned() + .collect::>() + .into_iter() + .rev() + .collect::>() + .join(" | "); + Err(WispError::Model(format!( + "download failed after {} retry rounds; partial data was kept for retry. {attempted}", + config.retries + ))) + } +} + +#[cfg(feature = "http")] +fn transfer_once( + agent: &ureq::Agent, + url: &str, + dest: &Path, + on_bytes: &mut dyn FnMut(u64), +) -> Result<()> { + let offset = dest.metadata().map(|metadata| metadata.len()).unwrap_or(0); + let mut request = agent.get(url).set("Accept-Encoding", "identity"); + if offset > 0 { + request = request.set("Range", &format!("bytes={offset}-")); + } + + let response = match request.call() { + Ok(response) => response, + Err(ureq::Error::Status(416, response)) if offset > 0 => { + let complete = response + .header("Content-Range") + .and_then(|value| value.rsplit_once('/')) + .and_then(|(_, total)| total.parse::().ok()) + == Some(offset); + if complete { + on_bytes(offset); + return Ok(()); + } + return Err(WispError::Model(format!( + "server rejected resume at byte {offset}" + ))); + } + Err(error) => return Err(WispError::Model(error.to_string())), + }; + + let append = offset > 0 && response.status() == 206; + if append { + let expected_prefix = format!("bytes {offset}-"); + if !response + .header("Content-Range") + .is_some_and(|value| value.starts_with(&expected_prefix)) + { + return Err(WispError::Model( + "server returned an invalid Content-Range".to_owned(), + )); + } + } + + let expected = response + .header("Content-Length") + .and_then(|value| value.parse::().ok()); + let mut reader = response.into_reader(); + let mut output = if append { + OpenOptions::new().create(true).append(true).open(dest)? + } else { + File::create(dest)? + }; + let start = if append { offset } else { 0 }; + on_bytes(start); + + let mut buffer = vec![0u8; 64 * 1024]; + let mut written = 0u64; + loop { + let count = reader.read(&mut buffer)?; + if count == 0 { + break; + } + output.write_all(&buffer[..count])?; + written += count as u64; + on_bytes(start + written); + } + output.flush()?; + + if let Some(total) = expected { + if total != written { + return Err(WispError::Model(format!( + "truncated response ({written} of {total} bytes)" + ))); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::net::{TcpListener, TcpStream}; + use std::thread; + + fn read_request(stream: &mut TcpStream) -> String { + let mut bytes = Vec::new(); + let mut buffer = [0u8; 1024]; + while !bytes.ends_with(b"\r\n\r\n") { + let count = stream.read(&mut buffer).unwrap(); + assert!(count > 0); + bytes.extend_from_slice(&buffer[..count]); + } + String::from_utf8(bytes).unwrap() + } + + fn direct_config(retries: u8) -> DownloadConfig { + DownloadConfig { + source: DownloadSource::Official, + proxy_mode: ProxyMode::Direct, + connect_timeout_secs: 2, + read_timeout_secs: 2, + retries, + ..DownloadConfig::default() + } + } + + #[test] + fn automatic_source_falls_back_between_official_and_mirror() { + let downloader = HttpDownloader::with_config(DownloadConfig::default()).unwrap(); + let official = "https://huggingface.co/org/model/resolve/main/model.onnx"; + let mirror = "https://hf-mirror.com/org/model/resolve/main/model.onnx"; + assert_eq!(downloader.candidate_urls(official), vec![official, mirror]); + + downloader.remember_success(official, mirror); + assert_eq!(downloader.candidate_urls(official), vec![mirror, official]); + } + + #[test] + fn mirror_first_uses_custom_endpoint_then_official() { + let config = DownloadConfig { + source: DownloadSource::MirrorFirst, + mirror_url: "https://models.example.cn/hf".to_owned(), + ..DownloadConfig::default() + }; + let downloader = HttpDownloader::with_config(config).unwrap(); + assert_eq!( + downloader.candidate_urls("https://huggingface.co/org/model/resolve/main/model.onnx"), + vec![ + "https://models.example.cn/hf/org/model/resolve/main/model.onnx", + "https://huggingface.co/org/model/resolve/main/model.onnx", + ] + ); + } + + #[test] + fn official_source_and_non_hugging_face_urls_are_not_rewritten() { + let config = DownloadConfig { + source: DownloadSource::Official, + ..DownloadConfig::default() + }; + let downloader = HttpDownloader::with_config(config).unwrap(); + assert_eq!( + downloader.candidate_urls("https://huggingface.co/org/model/resolve/main/file.bin"), + vec!["https://huggingface.co/org/model/resolve/main/file.bin"] + ); + assert_eq!( + downloader.candidate_urls("https://github.com/org/repo/releases/download/v1/file.bin"), + vec!["https://github.com/org/repo/releases/download/v1/file.bin"] + ); + } + + #[test] + fn cloned_downloaders_receive_live_setting_updates() { + let store_downloader = HttpDownloader::default(); + let settings_handle = store_downloader.clone(); + settings_handle + .update_config(DownloadConfig { + source: DownloadSource::MirrorFirst, + mirror_url: "https://models.example.cn".to_owned(), + ..DownloadConfig::default() + }) + .unwrap(); + + assert_eq!( + store_downloader + .candidate_urls("https://huggingface.co/org/model/resolve/main/model.onnx")[0], + "https://models.example.cn/org/model/resolve/main/model.onnx" + ); + } + + #[test] + fn settings_reject_insecure_mirror_and_incomplete_custom_proxy() { + let insecure = DownloadConfig { + mirror_url: "http://models.example.cn".to_owned(), + ..DownloadConfig::default() + }; + assert!(HttpDownloader::with_config(insecure).is_err()); + + let missing_proxy = DownloadConfig { + proxy_mode: ProxyMode::Custom, + proxy_url: String::new(), + ..DownloadConfig::default() + }; + assert!(HttpDownloader::with_config(missing_proxy).is_err()); + } + + #[test] + fn proxy_urls_cover_http_and_socks() { + for proxy_url in [ + "http://127.0.0.1:7890", + "socks5://127.0.0.1:1080", + "socks4a://proxy.example:1080", + ] { + let config = DownloadConfig { + proxy_mode: ProxyMode::Custom, + proxy_url: proxy_url.to_owned(), + ..DownloadConfig::default() + }; + assert!(HttpDownloader::with_config(config).is_ok(), "{proxy_url}"); + } + } + + #[test] + fn resumes_an_existing_partial_with_an_http_range() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let server = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + let request = read_request(&mut stream); + assert!(request.contains("Range: bytes=6-")); + stream + .write_all( + b"HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 6-10/11\r\nConnection: close\r\n\r\nworld", + ) + .unwrap(); + }); + + let temp = tempfile::tempdir().unwrap(); + let dest = temp.path().join("model.part"); + std::fs::write(&dest, b"hello ").unwrap(); + let downloader = HttpDownloader::with_config(direct_config(1)).unwrap(); + let mut progress = Vec::new(); + downloader + .download_with_progress(&format!("http://{address}/model"), &dest, &mut |bytes| { + progress.push(bytes) + }) + .unwrap(); + + server.join().unwrap(); + assert_eq!(std::fs::read(dest).unwrap(), b"hello world"); + assert_eq!(progress.last(), Some(&11)); + } + + #[test] + fn retries_a_dropped_transfer_from_its_last_byte() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let server = thread::spawn(move || { + let (mut first, _) = listener.accept().unwrap(); + let request = read_request(&mut first); + assert!(!request.contains("Range:")); + first + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Length: 11\r\nConnection: close\r\n\r\nhello ", + ) + .unwrap(); + drop(first); + + let (mut second, _) = listener.accept().unwrap(); + let request = read_request(&mut second); + assert!(request.contains("Range: bytes=6-")); + second + .write_all( + b"HTTP/1.1 206 Partial Content\r\nContent-Length: 5\r\nContent-Range: bytes 6-10/11\r\nConnection: close\r\n\r\nworld", + ) + .unwrap(); + }); + + let temp = tempfile::tempdir().unwrap(); + let dest = temp.path().join("model.part"); + let downloader = HttpDownloader::with_config(direct_config(2)).unwrap(); + downloader + .download_with_progress(&format!("http://{address}/model"), &dest, &mut |_| {}) + .unwrap(); + + server.join().unwrap(); + assert_eq!(std::fs::read(dest).unwrap(), b"hello world"); + } + + #[test] + #[ignore = "downloads a real model asset; run explicitly for release validation"] + fn real_hugging_face_mirror_download_smoke() { + let config = DownloadConfig { + source: DownloadSource::MirrorFirst, + proxy_mode: ProxyMode::Direct, + retries: 1, + ..DownloadConfig::default() + }; + let downloader = HttpDownloader::with_config(config).unwrap(); + let temp = tempfile::tempdir().unwrap(); + let dest = temp.path().join("tokens.txt.part"); + downloader + .download( + "https://huggingface.co/csukuangfj/sherpa-onnx-sense-voice-zh-en-ja-ko-yue-2024-07-17/resolve/main/tokens.txt", + &dest, + ) + .unwrap(); + assert!(dest.metadata().unwrap().len() > 300_000); } } diff --git a/crates/wisp-models/src/lib.rs b/crates/wisp-models/src/lib.rs index 7bf2ed6..e240d4d 100644 --- a/crates/wisp-models/src/lib.rs +++ b/crates/wisp-models/src/lib.rs @@ -18,9 +18,9 @@ pub mod store; pub use catalog::{builtin_catalog, denoise_models, diarization_models}; pub use cloud::cloud_catalog; pub use coreml::{coreml_asset, CoremlAsset}; -pub use download::FileDownloader; #[cfg(feature = "http")] pub use download::HttpDownloader; +pub use download::{DownloadConfig, DownloadSource, FileDownloader, ProxyMode, DEFAULT_HF_MIRROR}; pub use machine::{ family_runnable, model_fit, parse_meminfo_total_bytes, recommended_accurate_model, recommended_default_model, Accelerator, GpuTier, MachineProfile, ModelFit, diff --git a/crates/wisp-models/src/store.rs b/crates/wisp-models/src/store.rs index 102c2dd..5518b3f 100644 --- a/crates/wisp-models/src/store.rs +++ b/crates/wisp-models/src/store.rs @@ -18,8 +18,8 @@ const COMPLETE_MARKER: &str = ".wisp-ok"; /// file, at least its size). /// /// Downloads are atomic: each file lands in a `*.part` temporary, is verified, and only then -/// renamed into place — and a failed transfer removes its `*.part`, so an interrupted run never -/// leaves a partial or corrupt file that looks valid. +/// renamed into place. Interrupted transfers keep that non-final `*.part` so the HTTP downloader +/// can resume it on retry; a checksum/size failure still removes the corrupt partial. pub struct FsModelStore { root: PathBuf, catalog: Vec, @@ -102,13 +102,8 @@ impl FsModelStore { } let part = dir.join(format!("{}.part", file.name)); - if let Err(e) = self - .downloader - .download_with_progress(&file.url, &part, on_bytes) - { - let _ = fs::remove_file(&part); // don't leave a partial transfer behind - return Err(e); - } + self.downloader + .download_with_progress(&file.url, &part, on_bytes)?; if let Err(e) = self.verify(&part, file) { let _ = fs::remove_file(&part); @@ -181,15 +176,10 @@ impl FsModelStore { } let zip_part = dir.join(format!("{}.zip.part", asset.dir_name)); - if let Err(e) = - self.downloader - .download_with_progress(&asset.url, &zip_part, &mut |bytes| { - on_progress(bytes.min(total), total); - }) - { - let _ = fs::remove_file(&zip_part); - return Err(e); - } + self.downloader + .download_with_progress(&asset.url, &zip_part, &mut |bytes| { + on_progress(bytes.min(total), total); + })?; if let Err(e) = unzip_into(&zip_part, &dir) { let _ = fs::remove_file(&zip_part); @@ -462,7 +452,7 @@ mod tests { } #[test] - fn download_error_removes_the_part_file() { + fn download_error_preserves_the_part_file_for_resume() { let desc = single_file_descriptor("e", "https://example/e.bin", "e.bin", b"whatever"); let root = tempfile::tempdir().unwrap(); let store = FsModelStore::new(root.path(), vec![desc], Box::new(PartialThenError)); @@ -471,10 +461,7 @@ mod tests { assert!(matches!(err, WispError::Model(_))); let dir = root.path().join("e"); - assert!( - !dir.join("e.bin.part").exists(), - "a failed download must not leave its .part behind" - ); + assert_eq!(fs::read(dir.join("e.bin.part")).unwrap(), b"partial"); assert!(!dir.join(COMPLETE_MARKER).exists()); }