use super::*; use sha2::{Digest, Sha256}; use std::fs::{self, File, OpenOptions}; use std::io::{Read, Write}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::time::Duration; const DOWNLOAD_CONNECT_TIMEOUT: Duration = Duration::from_secs(10); const DOWNLOAD_STALL_TIMEOUT: Duration = Duration::from_secs(30); pub(crate) fn download_managed_artifact( id: ManagedArtifactId, models_path: &Path, cancel: &AtomicBool, verified_bytes: &AtomicU64, ) -> Result { download_artifact_with_cancel( id.model(), id.artifact(), models_path, cancel, verified_bytes, ) } pub(crate) fn validate_managed_artifact( id: ManagedArtifactId, models_path: &Path, cancel: &AtomicBool, verified_bytes: &AtomicU64, ) -> Result { let model = id.model(); let artifact = id.artifact(); let destination = artifact.path(model, models_path); let partial = artifact.partial_path(model, models_path); let (path, promote) = if destination.exists() { (destination.clone(), false) } else if partial.exists() { (partial.clone(), true) } else { return Err(format!("{} is not downloaded", artifact.label)); }; match verify(&path, artifact, model, cancel, verified_bytes) { Ok(DownloadOutcome::Stopped) => return Ok(DownloadOutcome::Stopped), Ok(DownloadOutcome::Complete) => {} Err(error) => { let marker = artifact.verification_path(model, models_path); if let Err(remove_error) = fs::remove_file(marker) && remove_error.kind() != std::io::ErrorKind::NotFound { return Err(format!( "{error}; could not remove checksum marker: {remove_error}" )); } return Err(error); } } if promote { fs::rename(partial, destination).map_err(|error| error.to_string())?; } mark_verified(model, artifact, models_path)?; Ok(DownloadOutcome::Complete) } pub(crate) fn delete_managed_artifact( id: ManagedArtifactId, models_path: &Path, ) -> Result<(), String> { let model = id.model(); let artifact = id.artifact(); for path in [ artifact.path(model, models_path), artifact.partial_path(model, models_path), artifact.verification_path(model, models_path), ] { match fs::remove_file(path) { Ok(()) => {} Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} Err(error) => return Err(error.to_string()), } } Ok(()) } #[cfg(test)] fn download_artifact( model: ModelChoice, artifact: &Artifact, models_path: &Path, ) -> Result { download_artifact_with_cancel( model, artifact, models_path, &AtomicBool::new(false), &AtomicU64::new(0), ) } fn download_artifact_with_cancel( model: ModelChoice, artifact: &Artifact, models_path: &Path, cancel: &AtomicBool, verified_bytes: &AtomicU64, ) -> Result { if cancel.load(Ordering::Relaxed) { return Ok(DownloadOutcome::Stopped); } let destination = artifact.path(model, models_path); if artifact.is_installed(model, models_path) { return Ok(DownloadOutcome::Complete); } if destination.exists() { if verify(&destination, artifact, model, cancel, verified_bytes)? == DownloadOutcome::Stopped { return Ok(DownloadOutcome::Stopped); } mark_verified(model, artifact, models_path)?; return Ok(DownloadOutcome::Complete); } let directory = destination .parent() .ok_or_else(|| "model artifact path has no parent directory".to_owned())?; fs::create_dir_all(directory).map_err(|error| error.to_string())?; let partial = artifact.partial_path(model, models_path); let partial_size = partial.metadata().map_or(0, |metadata| metadata.len()); if partial_size > artifact.size { File::create(&partial).map_err(|error| error.to_string())?; } if partial.metadata().map_or(0, |metadata| metadata.len()) != artifact.size && download_to_partial(artifact, &partial, cancel)? == DownloadOutcome::Stopped { return Ok(DownloadOutcome::Stopped); } if cancel.load(Ordering::Relaxed) { return Ok(DownloadOutcome::Stopped); } match verify(&partial, artifact, model, cancel, verified_bytes) { Ok(DownloadOutcome::Stopped) => return Ok(DownloadOutcome::Stopped), Ok(DownloadOutcome::Complete) => {} Err(error) => { fs::remove_file(&partial).map_err(|remove_error| { format!("{error}; could not remove partial file: {remove_error}") })?; return Err(error); } } fs::rename(partial, destination).map_err(|error| error.to_string())?; mark_verified(model, artifact, models_path)?; Ok(DownloadOutcome::Complete) } fn mark_verified( model: ModelChoice, artifact: &Artifact, models_path: &Path, ) -> Result<(), String> { fs::write( artifact.verification_path(model, models_path), artifact.sha256, ) .map_err(|error| error.to_string()) } fn verify( path: &Path, artifact: &Artifact, model: ModelChoice, cancel: &AtomicBool, verified_bytes: &AtomicU64, ) -> Result { verified_bytes.store(0, Ordering::Relaxed); let size = path.metadata().map_err(|error| error.to_string())?.len(); if size != artifact.size { return Err(format!( "{} has size {size}, expected {}", path.display(), artifact.size )); } let mut file = File::open(path).map_err(|error| error.to_string())?; let mut hasher = Sha256::new(); let mut buffer = vec![0; 1024 * 1024]; loop { if cancel.load(Ordering::Relaxed) { return Ok(DownloadOutcome::Stopped); } let count = file.read(&mut buffer).map_err(|error| error.to_string())?; if count == 0 { break; } hasher.update(&buffer[..count]); verified_bytes.fetch_add(count as u64, Ordering::Relaxed); } let actual = hex(&hasher.finalize()); if actual != artifact.sha256 { return Err(format!( "Checksum verification failed for {}", path.display() )); } if let Some(support) = artifact.support { crate::engine::validate_model_artifact(path, model, support)?; } Ok(DownloadOutcome::Complete) } fn hex(bytes: &[u8]) -> String { const DIGITS: &[u8; 16] = b"0123456789abcdef"; let mut result = String::with_capacity(bytes.len() * 2); for byte in bytes { result.push(DIGITS[(byte >> 4) as usize] as char); result.push(DIGITS[(byte & 0x0f) as usize] as char); } result } fn download_to_partial( artifact: &Artifact, partial: &Path, cancel: &AtomicBool, ) -> Result { download_url_to_partial(&artifact.url(), partial, cancel) } fn download_url_to_partial( url: &str, partial: &Path, cancel: &AtomicBool, ) -> Result { download_url_to_partial_with_timeout(url, partial, cancel, true, DOWNLOAD_STALL_TIMEOUT) } fn download_url_to_partial_with_timeout( url: &str, partial: &Path, cancel: &AtomicBool, https_only: bool, stall_timeout: Duration, ) -> Result { let offset = partial.metadata().map_or(0, |metadata| metadata.len()); let agent: ureq::Agent = ureq::Agent::config_builder() .https_only(https_only) .timeout_connect(Some(DOWNLOAD_CONNECT_TIMEOUT)) .timeout_recv_response(Some(stall_timeout)) .timeout_recv_body(Some(stall_timeout)) .build() .into(); let mut request = agent.get(url); if offset > 0 { request = request.header("Range", format!("bytes={offset}-")); } let mut response = request .call() .map_err(|error| format!("Model download failed: {error}"))?; let status = response.status().as_u16(); let append = offset > 0 && status == 206; if offset > 0 && status != 200 && status != 206 { return Err(format!( "Model server returned HTTP {status} while resuming at byte {offset}" )); } if append { validate_content_range(&response, offset)?; } let mut output = OpenOptions::new() .create(true) .write(true) .append(append) .truncate(!append) .open(partial) .map_err(|error| error.to_string())?; let mut body = response.body_mut().as_reader(); let mut buffer = vec![0; 1024 * 1024]; loop { if cancel.load(Ordering::Relaxed) { return Ok(DownloadOutcome::Stopped); } let count = body .read(&mut buffer) .map_err(|error| format!("Model download failed: {error}"))?; if count == 0 { break; } output .write_all(&buffer[..count]) .map_err(|error| error.to_string())?; } Ok(DownloadOutcome::Complete) } fn validate_content_range( response: &ureq::http::Response, offset: u64, ) -> Result<(), String> { let expected = format!("bytes {offset}-"); let content_range = response .headers() .get("content-range") .and_then(|value| value.to_str().ok()) .unwrap_or_default(); if content_range.starts_with(&expected) { Ok(()) } else { Err(format!( "Model server returned an invalid Content-Range while resuming at byte {offset}" )) } } #[cfg(test)] mod tests { use super::*; use std::net::TcpListener; use std::thread; use std::time::{SystemTime, UNIX_EPOCH}; #[test] fn catalog_and_checksum_verification_are_explicit() { assert_eq!(ModelChoice::from_id("glm-5.2"), Some(ModelChoice::Glm52)); assert_eq!( ModelChoice::from_id("deepseek-v4-flash-0731"), Some(ModelChoice::DeepSeekV4Flash0731) ); assert!(ModelChoice::from_id("unknown").is_none()); assert_eq!( ModelChoice::DeepSeekV4Flash.main_artifact().size, 86_720_111_488 ); assert_eq!(ModelChoice::Glm52.main_artifact().size, 211_075_856_448); assert_eq!( ModelChoice::DeepSeekV4Flash0731.main_artifact().size, 86_720_111_488 ); assert_eq!( ModelChoice::DeepSeekV4Flash.artifacts(true, true).count(), 3 ); assert_eq!(ModelChoice::Glm52.artifacts(true, true).count(), 1); assert_eq!( ModelChoice::DeepSeekV4Flash0731 .artifacts(true, true) .count(), 2 ); let id = SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos(); let models_path = std::env::temp_dir().join(format!("ds4-server-models-{id}")); let engine = engine_artifacts(ModelChoice::DeepSeekV4Flash, false, true, &models_path); assert_eq!( engine.model.file_name(), Some(std::ffi::OsStr::new(FLASH.file_name)) ); assert_eq!( engine.mtp.as_deref().and_then(Path::file_name), Some(std::ffi::OsStr::new(FLASH_DSPARK.file_name)) ); let flash_0731 = engine_artifacts(ModelChoice::DeepSeekV4Flash0731, false, true, &models_path); assert_eq!( flash_0731.mtp.as_deref().and_then(Path::file_name), Some(std::ffi::OsStr::new(FLASH_0731_DSPARK.file_name)) ); assert!( engine_artifacts(ModelChoice::DeepSeekV4Flash0731, true, false, &models_path) .mtp .is_none() ); let empty = Artifact { label: "empty", file_name: "empty", repository: "", size: 0, sha256: "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", support: None, }; let partial = empty.partial_path(ModelChoice::DeepSeekV4Flash, &models_path); fs::create_dir_all(partial.parent().unwrap()).unwrap(); fs::write(&partial, []).unwrap(); download_artifact(ModelChoice::DeepSeekV4Flash, &empty, &models_path).unwrap(); assert!(empty.is_installed(ModelChoice::DeepSeekV4Flash, &models_path)); assert!(!partial.exists()); assert_eq!( fs::read_to_string(empty.verification_path(ModelChoice::DeepSeekV4Flash, &models_path)) .unwrap(), empty.sha256 ); let installed = empty.path(ModelChoice::DeepSeekV4Flash, &models_path); assert!(empty.validate_installed_path(&installed).is_ok()); let wrong_name = installed.with_file_name("wrong-checkpoint.gguf"); fs::write(&wrong_name, []).unwrap(); fs::write(wrong_name.with_extension("gguf.sha256"), empty.sha256).unwrap(); assert!(empty.validate_installed_path(&wrong_name).is_err()); assert!( validate_engine_artifacts( ModelChoice::DeepSeekV4Flash0731, true, false, &EngineArtifacts { model: installed.clone(), mtp: None, }, ) .is_err() ); assert!( validate_engine_artifacts( ModelChoice::DeepSeekV4Pro, false, true, &EngineArtifacts { model: installed, mtp: None, }, ) .is_err() ); fs::remove_dir_all(models_path).unwrap(); } #[test] fn verification_reports_bytes_read() { let path = std::env::temp_dir().join(format!( "ds4-server-verify-{}", SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos() )); fs::write(&path, b"abc").unwrap(); let artifact = Artifact { label: "test model", file_name: "unused", repository: "unused", size: 3, sha256: "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad", support: None, }; let verified_bytes = AtomicU64::new(999); assert_eq!( verify( &path, &artifact, ModelChoice::DeepSeekV4Flash, &AtomicBool::new(false), &verified_bytes, ) .unwrap(), DownloadOutcome::Complete ); assert_eq!(verified_bytes.load(Ordering::Relaxed), 3); fs::remove_file(path).unwrap(); } #[test] fn managed_artifact_inventory_and_delete_include_partial_files() { let models_path = std::env::temp_dir().join(format!( "ds4-server-manager-{}", SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos() )); let id = ManagedArtifactId::DeepSeekV4Flash; let partial = id.artifact().partial_path(id.model(), &models_path); fs::create_dir_all(partial.parent().unwrap()).unwrap(); fs::write(&partial, b"part").unwrap(); let managed = managed_artifacts(&models_path) .into_iter() .find(|artifact| artifact.id == id) .unwrap(); assert_eq!(managed.stored, 4); assert_eq!(managed.state, ManagedArtifactState::Partial); delete_managed_artifact(id, &models_path).unwrap(); assert!(!partial.exists()); fs::remove_dir_all(models_path).unwrap(); } #[test] fn restart_resumes_at_the_existing_partial_byte() { let content = b"restart-resume works"; let offset = 8; let directory = std::env::temp_dir().join(format!( "ds4-server-resume-{}", SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos() )); fs::create_dir_all(&directory).unwrap(); let partial = directory.join("model.gguf.part"); fs::write(&partial, &content[..offset]).unwrap(); let listener = match TcpListener::bind("127.0.0.1:0") { Ok(listener) => listener, Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { fs::remove_dir_all(directory).unwrap(); return; } Err(error) => panic!("could not start test server: {error}"), }; let address = listener.local_addr().unwrap(); let server = thread::spawn(move || { let (mut connection, _) = listener.accept().unwrap(); let mut request = [0; 2048]; let count = connection.read(&mut request).unwrap(); let request = String::from_utf8_lossy(&request[..count]).to_ascii_lowercase(); assert!(request.contains("range: bytes=8-")); let remaining = &content[offset..]; write!( connection, "HTTP/1.1 206 Partial Content\r\nContent-Length: {}\r\nContent-Range: bytes {offset}-{}/{}\r\nConnection: close\r\n\r\n", remaining.len(), content.len() - 1, content.len(), ) .unwrap(); connection.write_all(remaining).unwrap(); }); let outcome = download_url_to_partial_with_timeout( &format!("http://{address}/model.gguf"), &partial, &AtomicBool::new(false), false, DOWNLOAD_STALL_TIMEOUT, ) .unwrap(); server.join().unwrap(); assert_eq!(outcome, DownloadOutcome::Complete); assert_eq!(fs::read(&partial).unwrap(), content); fs::remove_dir_all(directory).unwrap(); } #[test] fn stalled_download_times_out_and_keeps_the_partial_file() { let directory = std::env::temp_dir().join(format!( "ds4-server-stall-{}", SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos() )); fs::create_dir_all(&directory).unwrap(); let partial = directory.join("model.gguf.part"); fs::write(&partial, b"part").unwrap(); let listener = match TcpListener::bind("127.0.0.1:0") { Ok(listener) => listener, Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { fs::remove_dir_all(directory).unwrap(); return; } Err(error) => panic!("could not start test server: {error}"), }; let address = listener.local_addr().unwrap(); let server = thread::spawn(move || { let (mut connection, _) = listener.accept().unwrap(); let mut request = [0; 2048]; let _count = connection.read(&mut request).unwrap(); connection .write_all( b"HTTP/1.1 206 Partial Content\r\nContent-Length: 10\r\nContent-Range: bytes 4-13/14\r\nConnection: close\r\n\r\n", ) .unwrap(); thread::sleep(Duration::from_millis(200)); }); let error = download_url_to_partial_with_timeout( &format!("http://{address}/model.gguf"), &partial, &AtomicBool::new(false), false, Duration::from_millis(50), ) .unwrap_err(); server.join().unwrap(); assert!(error.contains("timeout"), "{error}"); assert_eq!(fs::read(&partial).unwrap(), b"part"); fs::remove_dir_all(directory).unwrap(); } #[test] fn cancellation_keeps_the_partial_file_for_the_next_run() { let directory = std::env::temp_dir().join(format!( "ds4-server-cancel-{}", SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap() .as_nanos() )); let artifact = Artifact { label: "test model", file_name: "model.gguf", repository: "unused", size: 10, sha256: "unused", support: None, }; let partial = artifact.partial_path(ModelChoice::DeepSeekV4Flash, &directory); fs::create_dir_all(partial.parent().unwrap()).unwrap(); fs::write(&partial, b"part").unwrap(); let cancel = AtomicBool::new(true); assert_eq!( download_artifact_with_cancel( ModelChoice::DeepSeekV4Flash, &artifact, &directory, &cancel, &AtomicU64::new(0), ) .unwrap(), DownloadOutcome::Stopped ); assert_eq!(fs::read(&partial).unwrap(), b"part"); fs::remove_dir_all(directory).unwrap(); } }