use std::net::IpAddr;
use crate::tools::ToolError;
pub fn is_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => {
let octets = v4.octets();
octets[0] == 127
|| octets[0] == 10
|| (octets[0] == 172 && octets[1] >= 16 && octets[1] <= 31)
|| (octets[0] == 192 && octets[1] == 168)
|| (octets[0] == 169 && octets[1] == 254)
|| *v4 == std::net::Ipv4Addr::UNSPECIFIED
}
IpAddr::V6(v6) => {
if let Some(v4) = v6.to_ipv4_mapped() {
return is_private_ip(&IpAddr::V4(v4));
}
v6.is_loopback()
|| (v6.segments()[0] & 0xfe00) == 0xfc00
|| matches!(v6.segments(), [0xfe80, ..])
|| *v6 == std::net::Ipv6Addr::UNSPECIFIED
}
}
}
pub async fn url_points_to_private_ip(url: &str) -> Result<bool, ToolError> {
let parsed =
url::Url::parse(url).map_err(|e| ToolError::InvalidInput(format!("Invalid URL: {}", e)))?;
let host = parsed
.host_str()
.ok_or_else(|| ToolError::InvalidInput("URL has no host".to_string()))?;
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(is_private_ip(&ip));
}
let port = parsed.port_or_known_default().unwrap_or(80);
let addr_str = format!("{}:{}", host, port);
let addrs: Vec<IpAddr> = tokio::net::lookup_host(&addr_str)
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("DNS resolution failed for {}: {}", host, e))
})?
.map(|sa| sa.ip())
.collect();
if addrs.is_empty() {
return Err(ToolError::ExecutionFailed(format!(
"DNS resolution returned no addresses for {}",
host
)));
}
Ok(addrs.iter().any(is_private_ip))
}
const MAX_REDIRECTS: usize = 10;
pub async fn guarded_get(
client: &reqwest::Client,
url: &str,
check_ssrf: bool,
) -> Result<reqwest::Response, ToolError> {
let mut current = url.to_string();
for _ in 0..=MAX_REDIRECTS {
if check_ssrf && url_points_to_private_ip(¤t).await? {
return Err(ToolError::ExecutionFailed(
"Request to private/internal IP address is blocked by SSRF protection. \
Call .with_allow_private_ips(true) to allow."
.to_string(),
));
}
let resp = client
.get(¤t)
.send()
.await
.map_err(|e| ToolError::ExecutionFailed(format!("HTTP request failed: {}", e)))?;
if !resp.status().is_redirection() {
return Ok(resp);
}
let Some(location) = resp
.headers()
.get(reqwest::header::LOCATION)
.and_then(|v| v.to_str().ok())
else {
return Ok(resp);
};
current = resolve_redirect(¤t, location)?;
}
Err(ToolError::ExecutionFailed(format!(
"request redirect count exceeded the limit of {} times",
MAX_REDIRECTS
)))
}
fn resolve_redirect(base: &str, location: &str) -> Result<String, ToolError> {
let joined = url::Url::parse(base)
.and_then(|base_url| base_url.join(location))
.map_err(|e| ToolError::InvalidInput(format!("invalid redirect target: {}", e)))?;
if joined.scheme() != "http" && joined.scheme() != "https" {
return Err(ToolError::InvalidInput(format!(
"redirect target protocol not supported: {}",
joined.scheme()
)));
}
Ok(joined.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn ipv4_mapped_ipv6_private_is_blocked() {
assert!(is_private_ip(
&"::ffff:127.0.0.1".parse::<IpAddr>().unwrap()
));
assert!(is_private_ip(&"::ffff:10.0.0.1".parse::<IpAddr>().unwrap()));
assert!(is_private_ip(
&"::ffff:169.254.169.254".parse::<IpAddr>().unwrap()
));
assert!(is_private_ip(
&"::ffff:192.168.1.1".parse::<IpAddr>().unwrap()
));
assert!(is_private_ip(
&"::ffff:172.16.0.1".parse::<IpAddr>().unwrap()
));
}
#[test]
fn ipv4_mapped_ipv6_public_allowed() {
assert!(!is_private_ip(&"::ffff:8.8.8.8".parse::<IpAddr>().unwrap()));
assert!(!is_private_ip(&"::ffff:1.1.1.1".parse::<IpAddr>().unwrap()));
}
#[test]
fn regular_ipv6_unchanged() {
assert!(is_private_ip(&"::1".parse::<IpAddr>().unwrap()));
assert!(is_private_ip(&"fc00::1".parse::<IpAddr>().unwrap()));
assert!(is_private_ip(&"fe80::1".parse::<IpAddr>().unwrap()));
assert!(!is_private_ip(&"2001:db8::1".parse::<IpAddr>().unwrap()));
}
#[test]
fn resolve_redirect_relative_and_absolute() {
assert_eq!(
resolve_redirect("https://a.com/x", "/internal").unwrap(),
"https://a.com/internal"
);
assert_eq!(
resolve_redirect("https://a.com/x", "https://b.com/y").unwrap(),
"https://b.com/y"
);
}
#[test]
fn resolve_redirect_rejects_non_http() {
assert!(resolve_redirect("https://a.com/x", "file:///etc/passwd").is_err());
assert!(resolve_redirect("https://a.com/x", "ftp://b.com").is_err());
}
}