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::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!(
|
||||||
|
|||||||
Reference in New Issue
Block a user