Skip to main content

apollo/tools/
network.rs

1use std::net::IpAddr;
2
3fn is_private_ip(ip: IpAddr) -> bool {
4    match ip {
5        IpAddr::V4(ipv4) => {
6            let octets = ipv4.octets();
7            octets[0] == 127
8                || octets[0] == 10
9                || (octets[0] == 172 && (16..=31).contains(&octets[1]))
10                || (octets[0] == 192 && octets[1] == 168)
11                || (octets[0] == 169 && octets[1] == 254)
12        }
13        IpAddr::V6(ipv6) => ipv6.is_loopback() || ((ipv6.segments()[0] & 0xfe00) == 0xfc00),
14    }
15}
16
17fn host_matches_allowlist(host: &str, allowed_domains: &[String]) -> bool {
18    allowed_domains.iter().any(|domain| {
19        host.eq_ignore_ascii_case(domain)
20            || host
21                .to_ascii_lowercase()
22                .ends_with(&format!(".{}", domain.to_ascii_lowercase()))
23    })
24}
25
26pub async fn validate_public_http_url(
27    url: &str,
28    allowed_domains: &[String],
29) -> anyhow::Result<reqwest::Url> {
30    let parsed = reqwest::Url::parse(url).map_err(|e| anyhow::anyhow!("Invalid URL: {}", e))?;
31    match parsed.scheme() {
32        "http" | "https" => {}
33        other => anyhow::bail!("Unsupported URL scheme: {}", other),
34    }
35
36    let host = parsed
37        .host_str()
38        .ok_or_else(|| anyhow::anyhow!("URL is missing a host"))?;
39
40    if host.eq_ignore_ascii_case("localhost")
41        || host.eq_ignore_ascii_case("0.0.0.0")
42        || host.ends_with(".localhost")
43    {
44        anyhow::bail!("Requests to local hosts are blocked");
45    }
46
47    if !allowed_domains.is_empty() && !host_matches_allowlist(host, allowed_domains) {
48        anyhow::bail!("Domain '{}' is not in the allowed list", host);
49    }
50
51    if let Ok(ip) = host.parse::<IpAddr>() {
52        if is_private_ip(ip) {
53            anyhow::bail!("Requests to private IP addresses are blocked");
54        }
55        return Ok(parsed);
56    }
57
58    let port = parsed.port_or_known_default().unwrap_or(80);
59    let resolved = tokio::net::lookup_host((host, port))
60        .await
61        .map_err(|e| anyhow::anyhow!("Failed to resolve host '{}': {}", host, e))?;
62
63    let mut saw_address = false;
64    for addr in resolved {
65        saw_address = true;
66        if is_private_ip(addr.ip()) {
67            anyhow::bail!("Host '{}' resolves to a private address", host);
68        }
69    }
70
71    if !saw_address {
72        anyhow::bail!("Host '{}' did not resolve to any addresses", host);
73    }
74
75    Ok(parsed)
76}
77
78#[cfg(test)]
79mod tests {
80    use super::*;
81
82    #[tokio::test]
83    async fn rejects_localhost() {
84        let err = validate_public_http_url("http://localhost:8080", &[])
85            .await
86            .unwrap_err();
87        assert!(err.to_string().contains("local hosts"));
88    }
89
90    #[tokio::test]
91    async fn rejects_private_ip() {
92        let err = validate_public_http_url("https://127.0.0.1", &[])
93            .await
94            .unwrap_err();
95        assert!(err.to_string().contains("private IP"));
96    }
97
98    #[tokio::test]
99    async fn rejects_bad_scheme() {
100        let err = validate_public_http_url("file:///etc/passwd", &[])
101            .await
102            .unwrap_err();
103        assert!(err.to_string().contains("Unsupported URL scheme"));
104    }
105}