diff --git a/src-tauri/src/captioner.rs b/src-tauri/src/captioner.rs index 37c6f79..b97da92 100644 --- a/src-tauri/src/captioner.rs +++ b/src-tauri/src/captioner.rs @@ -8,13 +8,20 @@ use ort::session::{builder::GraphOptimizationLevel, Session}; use ort::value::{Shape, Tensor}; use serde::{Deserialize, Serialize}; use std::borrow::Cow; -use std::io::{Cursor, Read}; +use std::io::Read; use std::path::{Path, PathBuf}; +use std::process::Command; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Mutex; use std::time::Instant; use tokenizers::Tokenizer; +// Suppress the console window when spawning curl.exe from the GUI app. +#[cfg(target_os = "windows")] +use std::os::windows::process::CommandExt; +#[cfg(target_os = "windows")] +const CREATE_NO_WINDOW: u32 = 0x08000000; + pub const FLORENCE_MODEL_ID: &str = "onnx-community/Florence-2-base-ft"; pub const FLORENCE_CAPTION_MODEL_NAME: &str = "florence-2-base-ft-onnx-q4"; const ONNX_RUNTIME_NUGET_URL: &str = @@ -723,50 +730,212 @@ fn download_onnx_runtime_files(local_dir: &Path) -> Result<()> { Ok(()) } +// Give up only after this many *consecutive* curl runs that download nothing; +// a run that makes any progress resets the counter, so a large file completes +// across however many resumes it takes. Kept low so a hard stall (e.g. a +// broken VM NIC) fails in a couple of minutes — surfacing a retryable error — +// rather than locking the UI on "preparing" for many minutes. +const MAX_STALL_RETRIES: usize = 3; + +/// Resiliently download `url` to `destination` using the system `curl.exe`. +/// +/// ureq's read timeout does not fire on a stalled large transfer on Windows +/// (schannel doesn't honor the socket read timeout), so a stall there hangs +/// forever. curl detects an inactivity stall (`--speed-time`), resumes from +/// the partial file (`-C -`), and retries internally — the same behavior a +/// browser gets. We monitor the `.part` file's size for the progress bar and +/// wrap curl in an outer progress-aware retry as a backstop. The partial file +/// survives an app restart, so a later retry continues from disk. +pub fn download_file_resilient( + url: &str, + destination: &Path, + mut on_progress: impl FnMut(u64, Option), +) -> Result<()> { + if let Some(parent) = destination.parent() { + std::fs::create_dir_all(parent)?; + } + let part = match destination.extension() { + Some(ext) => destination.with_extension(format!("{}.part", ext.to_string_lossy())), + None => destination.with_extension("part"), + }; + let name = destination + .file_name() + .map(|n| n.to_string_lossy().into_owned()) + .unwrap_or_else(|| url.to_string()); + + log::info!("{name}: resolving download size"); + let total = remote_content_length(url); + + // Reconcile any existing `.part` against the real size: exactly complete → + // finish; oversized (stale/corrupt) → discard so curl restarts cleanly + // (otherwise `curl -C -` would 416 forever). + if let Some(total) = total { + let size = std::fs::metadata(&part).map(|m| m.len()).unwrap_or(0); + if size == total { + std::fs::rename(&part, destination)?; + return Ok(()); + } + if size > total { + let _ = std::fs::remove_file(&part); + } + } + log::info!( + "{name}: downloading via curl ({} bytes)", + total + .map(|t| t.to_string()) + .unwrap_or_else(|| "unknown size".into()) + ); + + let mut stalls = 0usize; + loop { + let before = std::fs::metadata(&part).map(|m| m.len()).unwrap_or(0); + match run_curl_download(url, &part, total, &mut on_progress) { + Ok(()) => break, + Err(error) => { + let after = std::fs::metadata(&part).map(|m| m.len()).unwrap_or(0); + if after > before { + log::warn!("{name}: curl interrupted at {after} bytes, resuming: {error}"); + stalls = 0; + } else { + stalls += 1; + log::warn!( + "{name}: curl made no progress ({stalls}/{MAX_STALL_RETRIES}): {error}" + ); + if stalls >= MAX_STALL_RETRIES { + // Discard the partial so a future attempt restarts clean + // rather than getting stuck re-resuming a bad file. + let _ = std::fs::remove_file(&part); + return Err(error); + } + } + std::thread::sleep(std::time::Duration::from_secs(2)); + } + } + } + + if let Some(total) = total { + let got = std::fs::metadata(&part).map(|m| m.len()).unwrap_or(0); + if got < total { + anyhow::bail!("{name}: incomplete after curl ({got}/{total} bytes)"); + } + } + std::fs::rename(&part, destination)?; + Ok(()) +} + +/// Size probe via `curl -r 0-0` (a 1-byte Range request), parsing the total +/// from the `Content-Range: bytes 0-0/` header. Uses curl rather than +/// ureq so no part of the download path depends on ureq (which hangs on this +/// VM's TLS stack). Returns None if the server doesn't report a size. +fn remote_content_length(url: &str) -> Option { + let mut command = Command::new("curl.exe"); + command.args([ + "-sL", + "-r", + "0-0", + "-D", + "-", + "-o", + "NUL", + "--connect-timeout", + "30", + "--max-time", + "30", + url, + ]); + #[cfg(target_os = "windows")] + command.creation_flags(CREATE_NO_WINDOW); + let output = command.output().ok()?; + let headers = String::from_utf8_lossy(&output.stdout); + for line in headers.lines() { + if let Some(rest) = line.to_ascii_lowercase().strip_prefix("content-range:") { + if let Some(total) = rest.rsplit('/').next().map(str::trim) { + if let Ok(n) = total.parse::() { + return Some(n); + } + } + } + } + None +} + +/// Run one `curl.exe` download to `dest`, resuming from any partial file, while +/// reporting progress from the growing file size. Returns an error (leaving the +/// partial in place) if curl exits non-zero. +fn run_curl_download( + url: &str, + dest: &Path, + total: Option, + on_progress: &mut impl FnMut(u64, Option), +) -> Result<()> { + let mut command = Command::new("curl.exe"); + command + .arg("-fSL") // fail on HTTP errors, follow redirects, show errors + .args(["-C", "-"]) // resume from the existing output file + .args(["--retry", "3", "--retry-delay", "1", "--retry-connrefused"]) + .args(["--connect-timeout", "30"]) + // Abort (then --retry resumes) if under 1 KB/s for 30s — a real + // inactivity timeout, which is what ureq couldn't deliver here. + .args(["--speed-limit", "1024", "--speed-time", "30"]) + .arg("-s") // no progress meter (we watch the file instead) + .arg("-o") + .arg(dest) + .arg(url) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::piped()); + #[cfg(target_os = "windows")] + command.creation_flags(CREATE_NO_WINDOW); + + let mut child = command + .spawn() + .map_err(|e| anyhow::anyhow!("failed to launch curl.exe (required for downloads): {e}"))?; + + loop { + if let Some(status) = child.try_wait()? { + if status.success() { + return Ok(()); + } + let mut stderr = String::new(); + if let Some(mut pipe) = child.stderr.take() { + let _ = pipe.read_to_string(&mut stderr); + } + anyhow::bail!( + "curl exited with {}: {}", + status + .code() + .map(|c| c.to_string()) + .unwrap_or_else(|| "signal".into()), + stderr.trim() + ); + } + let downloaded = std::fs::metadata(dest).map(|m| m.len()).unwrap_or(0); + on_progress(downloaded, total); + std::thread::sleep(std::time::Duration::from_millis(300)); + } +} + fn download_nuget_file( source_url: &str, archive_path: &str, destination: &Path, - mut on_progress: impl FnMut(u64, Option), + on_progress: impl FnMut(u64, Option), ) -> Result<()> { - let mut response = ureq::get(source_url) - .call() - .map_err(|error| anyhow::anyhow!("{error}"))?; - let total_bytes = response - .headers() - .get("content-length") - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.parse::().ok()); + // Download the .nupkg (a zip) resiliently, then extract the one DLL. + let package = destination.with_extension("nupkg"); + download_file_resilient(source_url, &package, on_progress)?; - // Stream into memory (the zip reader needs Seek) while reporting bytes. - let mut bytes = Vec::with_capacity(total_bytes.unwrap_or(0) as usize); - let mut reader = response.body_mut().as_reader(); - let mut buffer = [0u8; 64 * 1024]; - let mut downloaded = 0u64; - loop { - let read = reader - .read(&mut buffer) - .map_err(|error| anyhow::anyhow!("{error}"))?; - if read == 0 { - break; - } - bytes.extend_from_slice(&buffer[..read]); - downloaded += read as u64; - on_progress(downloaded, total_bytes); - } - - let mut archive = zip::ZipArchive::new(Cursor::new(bytes))?; + log::info!("extracting {archive_path} from package"); + let file = std::fs::File::open(&package)?; + let mut archive = zip::ZipArchive::new(file)?; let mut dll = archive.by_name(archive_path)?; - if let Some(parent) = destination.parent() { - std::fs::create_dir_all(parent)?; - } - let temp_destination = destination.with_extension("tmp"); { - let mut file = std::fs::File::create(&temp_destination)?; - std::io::copy(&mut dll, &mut file)?; + let mut out = std::fs::File::create(&temp_destination)?; + std::io::copy(&mut dll, &mut out)?; } - std::fs::rename(temp_destination, destination)?; + std::fs::rename(&temp_destination, destination)?; + let _ = std::fs::remove_file(&package); + log::info!("extracted {archive_path}"); Ok(()) } diff --git a/src-tauri/src/tagger.rs b/src-tauri/src/tagger.rs index 5bc1a1f..a30f890 100644 --- a/src-tauri/src/tagger.rs +++ b/src-tauri/src/tagger.rs @@ -94,56 +94,6 @@ pub struct TaggerModelProgress { pub done: bool, } -/// Adapts hf-hub's `Progress` trait to throttled `tagger-model-progress` -/// events so the 1.3 GB model download reports real bytes, not a spinner. -struct HubProgress<'a, F: Fn(TaggerModelProgress)> { - emit: &'a F, - total_files: usize, - completed_files: usize, - current_file: String, - downloaded: u64, - total: u64, - last_emit: Instant, -} - -impl hf_hub::api::Progress for HubProgress<'_, F> { - fn init(&mut self, size: usize, _filename: &str) { - self.total = size as u64; - self.downloaded = 0; - self.last_emit = Instant::now(); - (self.emit)(self.snapshot(false)); - } - - fn update(&mut self, size: usize) { - self.downloaded = self.downloaded.saturating_add(size as u64); - if self.last_emit.elapsed() >= std::time::Duration::from_millis(200) { - self.last_emit = Instant::now(); - (self.emit)(self.snapshot(false)); - } - } - - fn finish(&mut self) { - (self.emit)(self.snapshot(false)); - } -} - -impl HubProgress<'_, F> { - fn snapshot(&self, done: bool) -> TaggerModelProgress { - TaggerModelProgress { - total_files: self.total_files, - completed_files: self.completed_files, - current_file: Some(self.current_file.clone()), - downloaded_bytes: Some(self.downloaded), - total_bytes: if self.total > 0 { - Some(self.total) - } else { - None - }, - done, - } - } -} - // --------------------------------------------------------------------------- // Runtime probe types exposed to the frontend // --------------------------------------------------------------------------- @@ -350,9 +300,14 @@ pub fn prepare_tagger_model_with_progress( )?; completed_files += dll_count; } + log::info!("Tagger: ONNX runtime DLLs ready; initializing runtime"); crate::captioner::ensure_onnx_runtime(&caption_model_dir)?; + log::info!("Tagger: runtime initialized; downloading model files"); - // ── Tagger model files (model.onnx is ~1.3 GB) ── + // ── Tagger model files (model.onnx is ~446 MB) ── + // Download directly from the resolved URL with our resilient downloader + // (timeout + resume), rather than hf-hub's download_with_progress, whose + // agent has no read timeout and would hang on a stalled connection. let api = Api::new()?; let repo = api.repo(Repo::new(WD_TAGGER_MODEL_ID.to_string(), RepoType::Model)); @@ -361,17 +316,22 @@ pub fn prepare_tagger_model_with_progress( if destination.exists() { continue; } - let progress = HubProgress { - emit: &emit_progress, - total_files, - completed_files, - current_file: (*file).to_string(), - downloaded: 0, - total: 0, - last_emit: Instant::now(), - }; - let cached = repo.download_with_progress(file, progress)?; - std::fs::copy(cached, destination)?; + let url = repo.url(file); + let label = (*file).to_string(); + let mut last_emit = Instant::now() - std::time::Duration::from_secs(1); + crate::captioner::download_file_resilient(&url, &destination, |downloaded, total| { + if last_emit.elapsed() >= std::time::Duration::from_millis(200) { + last_emit = Instant::now(); + emit_progress(TaggerModelProgress { + total_files, + completed_files, + current_file: Some(label.clone()), + downloaded_bytes: Some(downloaded), + total_bytes: total, + done: false, + }); + } + })?; completed_files += 1; }