Bound model download stalls

This commit is contained in:
Georg Bauer
2026-07-25 12:52:33 +02:00
parent 218c9473f7
commit 429cbc78de

View File

@@ -3,6 +3,10 @@ use sha2::{Digest, Sha256};
use std::fs::{self, File, OpenOptions}; use std::fs::{self, File, OpenOptions};
use std::io::{Read, Write}; use std::io::{Read, Write};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; 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( pub(crate) fn download_managed_artifact(
id: ManagedArtifactId, id: ManagedArtifactId,
@@ -229,10 +233,23 @@ fn download_url_to_partial(
url: &str, url: &str,
partial: &Path, partial: &Path,
cancel: &AtomicBool, cancel: &AtomicBool,
) -> Result<DownloadOutcome, String> {
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<DownloadOutcome, String> { ) -> Result<DownloadOutcome, String> {
let offset = partial.metadata().map_or(0, |metadata| metadata.len()); let offset = partial.metadata().map_or(0, |metadata| metadata.len());
let agent: ureq::Agent = ureq::Agent::config_builder() 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() .build()
.into(); .into();
let mut request = agent.get(url); let mut request = agent.get(url);
@@ -456,10 +473,12 @@ mod tests {
connection.write_all(remaining).unwrap(); connection.write_all(remaining).unwrap();
}); });
let outcome = download_url_to_partial( let outcome = download_url_to_partial_with_timeout(
&format!("http://{address}/model.gguf"), &format!("http://{address}/model.gguf"),
&partial, &partial,
&AtomicBool::new(false), &AtomicBool::new(false),
false,
DOWNLOAD_STALL_TIMEOUT,
) )
.unwrap(); .unwrap();
server.join().unwrap(); server.join().unwrap();
@@ -468,6 +487,55 @@ mod tests {
fs::remove_dir_all(directory).unwrap(); 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] #[test]
fn cancellation_keeps_the_partial_file_for_the_next_run() { fn cancellation_keeps_the_partial_file_for_the_next_run() {
let directory = std::env::temp_dir().join(format!( let directory = std::env::temp_dir().join(format!(