Skip to main content

studio_worker/
net.rs

1//! Shared transport-level guards for everything the worker downloads.
2//!
3//! Model files, `sd-cli` archives, ONNX runtimes, and auto-update
4//! installers all arrive over HTTP and are then either loaded into an
5//! engine or executed — so a plaintext-`http` fetch is a
6//! man-in-the-middle away from model poisoning or remote code
7//! execution.  Every downloader routes its URL through
8//! [`validate_download_url`] before the first byte is requested.
9
10use anyhow::{bail, Context, Result};
11
12/// Refuse any download URL that is not `https`.  Loopback `http` is
13/// allowed so test suites (wiremock) and air-gapped local mirrors keep
14/// working; everything else — remote `http`, `file`, `ftp`, garbage —
15/// is rejected before any request is made.  `what` names the artefact
16/// class (e.g. `"model file"`, `"installer"`) so the error tells the
17/// operator which config/registry entry to fix.
18pub fn validate_download_url(raw: &str, what: &str) -> Result<()> {
19    let url = url::Url::parse(raw).with_context(|| format!("invalid {what} URL {raw:?}"))?;
20    if url.scheme() == "https" {
21        return Ok(());
22    }
23    if url.scheme() == "http" {
24        // Typed hosts so bracketed IPv6 (`[::1]`) is recognised too —
25        // `host_str()` keeps the brackets and defeats `IpAddr::parse`.
26        match url.host() {
27            Some(url::Host::Domain(d)) if d.eq_ignore_ascii_case("localhost") => return Ok(()),
28            Some(url::Host::Ipv4(ip)) if ip.is_loopback() => return Ok(()),
29            Some(url::Host::Ipv6(ip)) if ip.is_loopback() => return Ok(()),
30            _ => {}
31        }
32    }
33    bail!("{what} URL must use https (loopback http is allowed for tests/mirrors): {raw}");
34}
35
36#[cfg(test)]
37mod tests {
38    use super::*;
39
40    #[test]
41    fn https_urls_pass() {
42        validate_download_url(
43            "https://huggingface.co/x/y/resolve/main/m.gguf",
44            "model file",
45        )
46        .unwrap();
47        validate_download_url(
48            "https://github.com/o/r/releases/download/v1/i.sh",
49            "installer",
50        )
51        .unwrap();
52    }
53
54    #[test]
55    fn loopback_http_passes_for_tests_and_mirrors() {
56        validate_download_url("http://127.0.0.1:1234/m.gguf", "model file").unwrap();
57        validate_download_url("http://localhost:1234/m.gguf", "model file").unwrap();
58        validate_download_url("http://[::1]:1234/m.gguf", "model file").unwrap();
59    }
60
61    #[test]
62    fn remote_http_is_rejected() {
63        let err = validate_download_url("http://example.com/m.gguf", "model file")
64            .unwrap_err()
65            .to_string();
66        assert!(err.contains("https"), "got: {err}");
67        assert!(
68            err.contains("model file"),
69            "must name the artefact class: {err}"
70        );
71    }
72
73    #[test]
74    fn non_http_schemes_are_rejected() {
75        for raw in [
76            "file:///etc/passwd",
77            "ftp://example.com/m.gguf",
78            "javascript:alert(1)",
79        ] {
80            let err = validate_download_url(raw, "model file")
81                .unwrap_err()
82                .to_string();
83            assert!(err.contains("https"), "{raw} must be rejected: {err}");
84        }
85    }
86
87    #[test]
88    fn malformed_urls_error_with_parse_context() {
89        let err = validate_download_url("not a url", "model file")
90            .unwrap_err()
91            .to_string();
92        assert!(err.contains("invalid model file URL"), "got: {err}");
93    }
94}