diff --git a/src/model/transfer.rs b/src/model/transfer.rs index 2854e6a..4f311e1 100644 --- a/src/model/transfer.rs +++ b/src/model/transfer.rs @@ -3,6 +3,10 @@ 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, @@ -229,10 +233,23 @@ 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(url.starts_with("https://")) + .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); @@ -456,10 +473,12 @@ mod tests { connection.write_all(remaining).unwrap(); }); - let outcome = download_url_to_partial( + 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(); @@ -468,6 +487,55 @@ mod tests { 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!(