1use anyhow::{bail, Context, Result};
11
12pub 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 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}