use anyhow::{bail, Context, Result};
pub fn validate_download_url(raw: &str, what: &str) -> Result<()> {
let url = url::Url::parse(raw).with_context(|| format!("invalid {what} URL {raw:?}"))?;
if url.scheme() == "https" {
return Ok(());
}
if url.scheme() == "http" {
match url.host() {
Some(url::Host::Domain(d)) if d.eq_ignore_ascii_case("localhost") => return Ok(()),
Some(url::Host::Ipv4(ip)) if ip.is_loopback() => return Ok(()),
Some(url::Host::Ipv6(ip)) if ip.is_loopback() => return Ok(()),
_ => {}
}
}
bail!("{what} URL must use https (loopback http is allowed for tests/mirrors): {raw}");
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn https_urls_pass() {
validate_download_url(
"https://huggingface.co/x/y/resolve/main/m.gguf",
"model file",
)
.unwrap();
validate_download_url(
"https://github.com/o/r/releases/download/v1/i.sh",
"installer",
)
.unwrap();
}
#[test]
fn loopback_http_passes_for_tests_and_mirrors() {
validate_download_url("http://127.0.0.1:1234/m.gguf", "model file").unwrap();
validate_download_url("http://localhost:1234/m.gguf", "model file").unwrap();
validate_download_url("http://[::1]:1234/m.gguf", "model file").unwrap();
}
#[test]
fn remote_http_is_rejected() {
let err = validate_download_url("http://example.com/m.gguf", "model file")
.unwrap_err()
.to_string();
assert!(err.contains("https"), "got: {err}");
assert!(
err.contains("model file"),
"must name the artefact class: {err}"
);
}
#[test]
fn non_http_schemes_are_rejected() {
for raw in [
"file:///etc/passwd",
"ftp://example.com/m.gguf",
"javascript:alert(1)",
] {
let err = validate_download_url(raw, "model file")
.unwrap_err()
.to_string();
assert!(err.contains("https"), "{raw} must be rejected: {err}");
}
}
#[test]
fn malformed_urls_error_with_parse_context() {
let err = validate_download_url("not a url", "model file")
.unwrap_err()
.to_string();
assert!(err.contains("invalid model file URL"), "got: {err}");
}
}