Bound model download stalls
This commit is contained in:
@@ -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<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> {
|
||||
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!(
|
||||
|
||||
Reference in New Issue
Block a user