Skip to main content

koan_core/remote/
download.rs

1//! Streaming file downloads: temp file, progress reporting, atomic rename, retries.
2//!
3//! Every remote byte koan writes to disk goes through here. `dest` only ever
4//! appears once the transfer completed, so a partially-written file can never
5//! be mistaken for a cached track.
6
7use std::io::{Read, Write};
8use std::path::Path;
9use std::time::Duration;
10
11use thiserror::Error;
12
13/// Longest the TCP connect + TLS handshake may take.
14const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
15
16/// Longest a single body read may block before the transfer counts as stalled.
17///
18/// `reqwest`'s blocking client re-applies its request timeout to each `Read`
19/// of a streamed response, so this bounds *stalls*, not total transfer time —
20/// a large file on a slow link keeps going as long as bytes keep arriving.
21const STALL_TIMEOUT: Duration = Duration::from_secs(30);
22
23/// Total deadline for JSON API calls, whose bodies are small and read in one go.
24pub const API_TIMEOUT: Duration = Duration::from_secs(30);
25
26/// Attempts a download gets before giving up.
27pub const DEFAULT_ATTEMPTS: u32 = 3;
28
29/// Base backoff between attempts; doubles each retry.
30const 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    /// Whether another attempt could plausibly succeed: transport-level
48    /// failures, truncated bodies, and server-side/rate-limit statuses.
49    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
61/// HTTP client for streaming large bodies — bounded connect, bounded stalls,
62/// no total deadline on the transfer.
63pub 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
70/// HTTP client for small JSON API calls, where a total request deadline is correct.
71pub 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
78/// Download to `dest`, retrying transient failures with exponential backoff.
79///
80/// `request` is invoked once per attempt so per-request state (Subsonic auth
81/// salts, for one) is regenerated rather than replayed; a request that cannot
82/// be built is fatal, not retried. `on_progress` receives
83/// `(bytes_this_attempt, total)` where `total` is 0 if the server sent no
84/// Content-Length; it restarts from zero when an attempt is retried.
85///
86/// Returns the number of bytes written. `dest` is left untouched on failure.
87pub 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    // Only once every attempt is spent. Between attempts the `.part` stays, and
124    // the next one truncates it in place: a stream already reading it holds that
125    // inode, and picks up again as the retry rewrites the same bytes.
126    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    // Status first: a 503 with a JSON body is transient and worth retrying.
138    if !status.is_success() {
139        return Err(DownloadError::Status(status));
140    }
141    // Subsonic reports failure with HTTP 200 and a JSON or XML error body, so a
142    // success status proves nothing on a binary endpoint. Without this, an error
143    // response gets written to disk and cached as if it were audio — it then
144    // reports Ready and fails to decode forever.
145    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
158/// Stream a response body into `dest` via a `.part` sibling, renaming only once
159/// the transfer completes. A read error or a body shorter than the advertised
160/// Content-Length errors, leaving the temp file for the caller to retry into.
161fn 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
212/// The in-progress sibling of `dest`. Appends `.part` rather than replacing the
213/// extension, so `Song.flac` and `Song.mp3` never collide on one temp file.
214pub 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
220/// Strip a `.part` suffix, yielding the final path a download will land at.
221/// Returns `path` unchanged when it isn't a temp file.
222pub 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    /// How a stub server answers one request.
241    #[derive(Clone)]
242    enum Reply {
243        /// Content-Length header, then that many bytes.
244        Complete(Vec<u8>),
245        /// Content-Length claims `claimed` bytes but only `body` is sent, then close.
246        Truncated {
247            claimed: usize,
248            body: Vec<u8>,
249        },
250        /// Chunked with no Content-Length, cut off mid-stream — what Navidrome
251        /// does for transcoded streams when the connection drops.
252        ChunkedTruncated(Vec<u8>),
253        ServerError,
254    }
255
256    /// Single-threaded stub HTTP server. Serves `replies` in order, repeating
257    /// the last one forever. Shuts down when the returned handle is dropped.
258    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                            // BSD sockets inherit O_NONBLOCK from the listener.
279                            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        // Drain the request headers so the client isn't left writing into a
316        // closed socket before it can read the response.
317        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                // No terminating zero-length chunk — the stream just stops.
352            }
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        // No Content-Length at all — the only signal is the stream ending
405        // mid-message, which must not be read as a finished download.
406        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    /// A stream reading the `.part` holds its inode, so a retry has to write
443    /// into that same file rather than a new one beside it.
444    #[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}