1use std::io::{Read, Write};
8use std::path::Path;
9use std::time::Duration;
10
11use thiserror::Error;
12
13const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
15
16const STALL_TIMEOUT: Duration = Duration::from_secs(30);
22
23pub const API_TIMEOUT: Duration = Duration::from_secs(30);
25
26pub const DEFAULT_ATTEMPTS: u32 = 3;
28
29const BACKOFF_BASE: Duration = Duration::from_millis(500);
31
32#[derive(Debug, Error)]
33pub enum DownloadError {
34 #[error("http error: {0}")]
35 Http(#[from] reqwest::Error),
36 #[error("io error: {0}")]
37 Io(#[from] std::io::Error),
38 #[error("incomplete download: got {got} of {expected} bytes")]
39 Incomplete { got: u64, expected: u64 },
40 #[error("server returned {0}")]
41 Status(reqwest::StatusCode),
42 #[error("request could not be built: {0}")]
43 Request(String),
44}
45
46impl DownloadError {
47 pub fn is_retryable(&self) -> bool {
50 match self {
51 DownloadError::Http(e) => e.is_timeout() || e.is_connect() || e.is_request(),
52 DownloadError::Io(_) | DownloadError::Incomplete { .. } => true,
53 DownloadError::Status(s) => {
54 s.is_server_error() || *s == reqwest::StatusCode::TOO_MANY_REQUESTS
55 }
56 DownloadError::Request(_) => false,
57 }
58 }
59}
60
61pub fn download_client() -> reqwest::Result<reqwest::blocking::Client> {
64 reqwest::blocking::Client::builder()
65 .connect_timeout(CONNECT_TIMEOUT)
66 .timeout(STALL_TIMEOUT)
67 .build()
68}
69
70pub fn api_client() -> reqwest::Result<reqwest::blocking::Client> {
72 reqwest::blocking::Client::builder()
73 .connect_timeout(CONNECT_TIMEOUT)
74 .timeout(API_TIMEOUT)
75 .build()
76}
77
78pub fn download_with_retries(
88 dest: &Path,
89 attempts: u32,
90 request: impl Fn() -> Result<reqwest::blocking::RequestBuilder, DownloadError>,
91 on_progress: impl Fn(u64, u64),
92) -> Result<u64, DownloadError> {
93 let attempts = attempts.max(1);
94 let mut last_err = None;
95
96 for attempt in 0..attempts {
97 if attempt > 0 {
98 let backoff = BACKOFF_BASE * 2u32.pow(attempt - 1);
99 log::warn!(
100 "download of {} failed ({}), retrying in {:?} ({}/{})",
101 dest.display(),
102 last_err
103 .as_ref()
104 .map(|e: &DownloadError| e.to_string())
105 .unwrap_or_default(),
106 backoff,
107 attempt + 1,
108 attempts
109 );
110 std::thread::sleep(backoff);
111 }
112
113 match attempt_download(dest, &request, &on_progress) {
114 Ok(bytes) => return Ok(bytes),
115 Err(e) if e.is_retryable() && attempt + 1 < attempts => last_err = Some(e),
116 Err(e) => {
117 last_err = Some(e);
118 break;
119 }
120 }
121 }
122
123 let _ = std::fs::remove_file(part_path(dest));
127 Err(last_err.expect("loop runs at least once and only exits here on error"))
128}
129
130fn attempt_download(
131 dest: &Path,
132 request: &impl Fn() -> Result<reqwest::blocking::RequestBuilder, DownloadError>,
133 on_progress: &impl Fn(u64, u64),
134) -> Result<u64, DownloadError> {
135 let resp = request()?.send()?;
136 let status = resp.status();
137 if !status.is_success() {
139 return Err(DownloadError::Status(status));
140 }
141 if resp
146 .headers()
147 .get(reqwest::header::CONTENT_TYPE)
148 .and_then(|v| v.to_str().ok())
149 .is_some_and(|ct| ct.contains("json") || ct.contains("xml"))
150 {
151 return Err(DownloadError::Request(
152 "server returned an error document where audio was expected".into(),
153 ));
154 }
155 stream_to_file(resp, dest, on_progress)
156}
157
158fn stream_to_file(
162 mut resp: reqwest::blocking::Response,
163 dest: &Path,
164 on_progress: &impl Fn(u64, u64),
165) -> Result<u64, DownloadError> {
166 let total = resp.content_length().unwrap_or(0);
167
168 if let Some(parent) = dest.parent() {
169 std::fs::create_dir_all(parent)?;
170 }
171
172 let tmp = part_path(dest);
173 let mut file = std::fs::File::create(&tmp)?;
174 let mut downloaded: u64 = 0;
175 let mut buf = [0u8; 64 * 1024];
176
177 let result = loop {
178 match resp.read(&mut buf) {
179 Ok(0) => break Ok(()),
180 Ok(n) => {
181 if let Err(e) = file.write_all(&buf[..n]) {
182 break Err(DownloadError::Io(e));
183 }
184 downloaded += n as u64;
185 on_progress(downloaded, total);
186 }
187 Err(e) => break Err(DownloadError::Io(e)),
188 }
189 };
190
191 let flushed = file.flush();
192 drop(file);
193
194 let outcome = result
195 .and_then(|()| flushed.map_err(DownloadError::Io))
196 .and_then(|()| {
197 if total > 0 && downloaded != total {
198 Err(DownloadError::Incomplete {
199 got: downloaded,
200 expected: total,
201 })
202 } else {
203 Ok(())
204 }
205 });
206
207 outcome?;
208 std::fs::rename(&tmp, dest)?;
209 Ok(downloaded)
210}
211
212pub fn part_path(dest: &Path) -> std::path::PathBuf {
215 let mut name = dest.file_name().unwrap_or_default().to_os_string();
216 name.push(".part");
217 dest.with_file_name(name)
218}
219
220pub fn strip_part_suffix(path: &Path) -> std::path::PathBuf {
223 let Some(name) = path.file_name().and_then(|n| n.to_str()) else {
224 return path.to_path_buf();
225 };
226 match name.strip_suffix(".part") {
227 Some(stripped) => path.with_file_name(stripped),
228 None => path.to_path_buf(),
229 }
230}
231
232#[cfg(test)]
233mod tests {
234 use super::*;
235 use std::io::BufRead;
236 use std::net::{TcpListener, TcpStream};
237 use std::sync::Arc;
238 use std::sync::atomic::{AtomicUsize, Ordering};
239
240 #[derive(Clone)]
242 enum Reply {
243 Complete(Vec<u8>),
245 Truncated {
247 claimed: usize,
248 body: Vec<u8>,
249 },
250 ChunkedTruncated(Vec<u8>),
253 ServerError,
254 }
255
256 struct StubServer {
259 addr: std::net::SocketAddr,
260 hits: Arc<AtomicUsize>,
261 shutdown: Arc<std::sync::atomic::AtomicBool>,
262 }
263
264 impl StubServer {
265 fn start(replies: Vec<Reply>) -> Self {
266 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
267 listener.set_nonblocking(true).unwrap();
268 let addr = listener.local_addr().unwrap();
269 let hits = Arc::new(AtomicUsize::new(0));
270 let shutdown = Arc::new(std::sync::atomic::AtomicBool::new(false));
271
272 let hits_bg = hits.clone();
273 let shutdown_bg = shutdown.clone();
274 std::thread::spawn(move || {
275 while !shutdown_bg.load(Ordering::Relaxed) {
276 match listener.accept() {
277 Ok((stream, _)) => {
278 let _ = stream.set_nonblocking(false);
280 let n = hits_bg.fetch_add(1, Ordering::SeqCst);
281 let reply = replies[n.min(replies.len() - 1)].clone();
282 serve_one(stream, reply);
283 }
284 Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
285 std::thread::sleep(Duration::from_millis(5));
286 }
287 Err(_) => break,
288 }
289 }
290 });
291
292 Self {
293 addr,
294 hits,
295 shutdown,
296 }
297 }
298
299 fn url(&self) -> String {
300 format!("http://{}/file", self.addr)
301 }
302
303 fn hits(&self) -> usize {
304 self.hits.load(Ordering::SeqCst)
305 }
306 }
307
308 impl Drop for StubServer {
309 fn drop(&mut self) {
310 self.shutdown.store(true, Ordering::Relaxed);
311 }
312 }
313
314 fn serve_one(mut stream: TcpStream, reply: Reply) {
315 let mut reader = std::io::BufReader::new(stream.try_clone().unwrap());
318 let mut line = String::new();
319 while reader.read_line(&mut line).unwrap_or(0) > 0 {
320 if line == "\r\n" || line == "\n" {
321 break;
322 }
323 line.clear();
324 }
325
326 match reply {
327 Reply::Complete(body) => {
328 let _ = write!(
329 stream,
330 "HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n",
331 body.len()
332 );
333 let _ = stream.write_all(&body);
334 }
335 Reply::Truncated { claimed, body } => {
336 let _ = write!(
337 stream,
338 "HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n",
339 claimed
340 );
341 let _ = stream.write_all(&body);
342 }
343 Reply::ChunkedTruncated(body) => {
344 let _ = write!(
345 stream,
346 "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n"
347 );
348 let _ = write!(stream, "{:x}\r\n", body.len());
349 let _ = stream.write_all(&body);
350 let _ = stream.write_all(b"\r\n");
351 }
353 Reply::ServerError => {
354 let _ = write!(stream, "HTTP/1.1 500 Internal Server Error\r\n\r\n");
355 }
356 }
357 let _ = stream.flush();
358 let _ = stream.shutdown(std::net::Shutdown::Both);
359 }
360
361 fn tmp_dest(dir: &tempfile::TempDir) -> std::path::PathBuf {
362 dir.path().join("nested").join("track.flac")
363 }
364
365 #[test]
366 fn complete_download_lands_at_dest() {
367 let body = vec![7u8; 200_000];
368 let server = StubServer::start(vec![Reply::Complete(body.clone())]);
369 let dir = tempfile::tempdir().unwrap();
370 let dest = tmp_dest(&dir);
371 let client = download_client().unwrap();
372
373 let written =
374 download_with_retries(&dest, 1, || Ok(client.get(server.url())), |_, _| {}).unwrap();
375
376 assert_eq!(written, body.len() as u64);
377 assert_eq!(std::fs::read(&dest).unwrap(), body);
378 assert!(!part_path(&dest).exists(), "temp file should be cleaned up");
379 }
380
381 #[test]
382 fn truncated_body_errors_and_leaves_no_file() {
383 let server = StubServer::start(vec![Reply::Truncated {
384 claimed: 100_000,
385 body: vec![1u8; 4_096],
386 }]);
387 let dir = tempfile::tempdir().unwrap();
388 let dest = tmp_dest(&dir);
389 let client = download_client().unwrap();
390
391 let err = download_with_retries(&dest, 1, || Ok(client.get(server.url())), |_, _| {})
392 .expect_err("a short body must not succeed");
393
394 assert!(
395 matches!(err, DownloadError::Incomplete { .. } | DownloadError::Io(_)),
396 "unexpected error: {err}"
397 );
398 assert!(!dest.exists(), "dest must not hold a truncated file");
399 assert!(!part_path(&dest).exists(), "temp file must be removed");
400 }
401
402 #[test]
403 fn missing_content_length_truncation_errors_rather_than_completing() {
404 let server = StubServer::start(vec![Reply::ChunkedTruncated(vec![9u8; 8_192])]);
407 let dir = tempfile::tempdir().unwrap();
408 let dest = tmp_dest(&dir);
409 let client = download_client().unwrap();
410
411 let err = download_with_retries(&dest, 1, || Ok(client.get(server.url())), |_, _| {})
412 .expect_err("a cut-off chunked body must not succeed");
413
414 assert!(matches!(err, DownloadError::Io(_)), "unexpected: {err}");
415 assert!(!dest.exists(), "dest must not hold a truncated file");
416 assert!(!part_path(&dest).exists(), "temp file must be removed");
417 }
418
419 #[test]
420 fn retries_transient_failure_then_succeeds() {
421 let body = vec![3u8; 50_000];
422 let server = StubServer::start(vec![
423 Reply::ServerError,
424 Reply::Truncated {
425 claimed: 50_000,
426 body: vec![3u8; 10],
427 },
428 Reply::Complete(body.clone()),
429 ]);
430 let dir = tempfile::tempdir().unwrap();
431 let dest = tmp_dest(&dir);
432 let client = download_client().unwrap();
433
434 let written =
435 download_with_retries(&dest, 3, || Ok(client.get(server.url())), |_, _| {}).unwrap();
436
437 assert_eq!(written, body.len() as u64);
438 assert_eq!(server.hits(), 3, "should have used all three attempts");
439 assert_eq!(std::fs::read(&dest).unwrap(), body);
440 }
441
442 #[cfg(unix)]
445 #[test]
446 fn retry_rewrites_the_same_part_file() {
447 use std::os::unix::fs::MetadataExt;
448
449 let server = StubServer::start(vec![
450 Reply::Truncated {
451 claimed: 50_000,
452 body: vec![3u8; 10],
453 },
454 Reply::Complete(vec![3u8; 50_000]),
455 ]);
456 let dir = tempfile::tempdir().unwrap();
457 let dest = tmp_dest(&dir);
458 let client = download_client().unwrap();
459
460 let inodes = std::sync::Mutex::new(std::collections::HashSet::new());
461 download_with_retries(
462 &dest,
463 2,
464 || Ok(client.get(server.url())),
465 |_, _| {
466 if let Ok(meta) = std::fs::metadata(part_path(&dest)) {
467 inodes.lock().unwrap().insert(meta.ino());
468 }
469 },
470 )
471 .unwrap();
472
473 assert_eq!(inodes.lock().unwrap().len(), 1);
474 }
475
476 #[test]
477 fn progress_reports_total_when_content_length_present() {
478 let body = vec![0u8; 300_000];
479 let server = StubServer::start(vec![Reply::Complete(body.clone())]);
480 let dir = tempfile::tempdir().unwrap();
481 let dest = tmp_dest(&dir);
482 let client = download_client().unwrap();
483
484 let seen = std::sync::Mutex::new(Vec::new());
485 download_with_retries(
486 &dest,
487 1,
488 || Ok(client.get(server.url())),
489 |d, t| {
490 seen.lock().unwrap().push((d, t));
491 },
492 )
493 .unwrap();
494
495 let seen = seen.into_inner().unwrap();
496 assert!(!seen.is_empty(), "progress should be reported");
497 assert!(seen.iter().all(|(_, t)| *t == body.len() as u64));
498 assert_eq!(seen.last().unwrap().0, body.len() as u64);
499 }
500
501 #[test]
502 fn part_path_appends_rather_than_replacing_extension() {
503 let flac = part_path(Path::new("/tmp/Song.flac"));
504 let mp3 = part_path(Path::new("/tmp/Song.mp3"));
505 assert_eq!(flac, Path::new("/tmp/Song.flac.part"));
506 assert_ne!(flac, mp3, "different codecs must not share a temp file");
507 }
508
509 #[test]
510 fn strip_part_suffix_round_trips() {
511 let dest = Path::new("/tmp/a/Song.flac");
512 assert_eq!(strip_part_suffix(&part_path(dest)), dest);
513 assert_eq!(strip_part_suffix(dest), dest);
514 }
515}