use std::net::{IpAddr, SocketAddr};
use std::time::Duration;
use crate::tools::ToolError;
pub fn is_private_ip(ip: &IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => {
if v4.octets()[0] == 0 {
return true;
}
is_private_ipv4(*v4)
}
IpAddr::V6(v6) => {
if let Some(v4) = v6.to_ipv4_mapped() {
return is_private_ip(&IpAddr::V4(v4));
}
is_private_ipv6(*v6)
}
}
}
fn is_private_ipv4(v4: std::net::Ipv4Addr) -> bool {
let [a, b, _, _] = v4.octets();
match a {
0 => true,
10 => true,
100 => b & 0b1100_0000 == 0b0100_0000,
127 => true,
169 => b == 254,
172 => (16..=31).contains(&b),
192 if b == 0 => true,
192 if b == 88 => true,
192 if b == 168 => true,
198 => (b & 0xfe) == 0x12 || b == 51,
203 if b == 0 => v4.octets()[2] == 113,
224..=255 => true,
_ => false,
}
}
fn is_private_ipv6(v6: std::net::Ipv6Addr) -> bool {
let seg = v6.segments();
if v6.is_loopback() {
return true;
}
if (seg[0] & 0xfe00) == 0xfc00 {
return true;
}
if matches!(seg, [0xfe80, ..]) {
return true;
}
if v6 == std::net::Ipv6Addr::UNSPECIFIED {
return true;
}
if (seg[0] & 0xffc0) == 0xfec0 {
return true;
}
if (seg[0] & 0xff00) == 0xff00 {
return true;
}
if seg[0] == 0x2001 && seg[1] == 0x0db8 {
return true;
}
seg[0] == 0x2002
}
pub const DEFAULT_GUARDED_TIMEOUT: Duration = Duration::from_secs(30);
pub async fn url_points_to_private_ip(url: &str) -> Result<bool, ToolError> {
let parsed = parse_http_url(url)?;
let addrs = resolve_url_addrs(&parsed).await?;
Ok(addrs.iter().any(|sa| is_private_ip(&sa.ip())))
}
fn parse_http_url(url: &str) -> Result<url::Url, ToolError> {
let parsed =
url::Url::parse(url).map_err(|e| ToolError::InvalidInput(format!("Invalid URL: {}", e)))?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err(ToolError::InvalidInput(format!(
"URL scheme not supported: {}",
parsed.scheme()
)));
}
if parsed.host_str().is_none() {
return Err(ToolError::InvalidInput("URL has no host".to_string()));
}
Ok(parsed)
}
async fn resolve_url_addrs(parsed: &url::Url) -> Result<Vec<SocketAddr>, ToolError> {
let host = parsed
.host_str()
.ok_or_else(|| ToolError::InvalidInput("URL has no host".to_string()))?;
let port = parsed.port_or_known_default().unwrap_or(80);
if let Ok(ip) = host.parse::<IpAddr>() {
return Ok(vec![SocketAddr::new(ip, port)]);
}
let addrs: Vec<SocketAddr> = tokio::net::lookup_host((host, port))
.await
.map_err(|e| {
ToolError::ExecutionFailed(format!("DNS resolution failed for {}: {}", host, e))
})?
.collect();
if addrs.is_empty() {
return Err(ToolError::ExecutionFailed(format!(
"DNS resolution returned no addresses for {}",
host
)));
}
Ok(addrs)
}
fn ensure_all_public(addrs: &[SocketAddr]) -> Result<(), ToolError> {
if let Some(bad) = addrs.iter().map(SocketAddr::ip).find(is_private_ip) {
return Err(ToolError::ExecutionFailed(format!(
"Request to private/internal IP address ({bad}) is blocked by SSRF protection. \
Call .with_allow_private_ips(true) to allow."
)));
}
Ok(())
}
fn pinned_client(
host: &str,
addrs: &[SocketAddr],
timeout: Option<Duration>,
) -> Result<reqwest::Client, ToolError> {
let mut builder = reqwest::Client::builder().redirect(reqwest::redirect::Policy::none());
builder = builder.timeout(timeout.unwrap_or(DEFAULT_GUARDED_TIMEOUT));
if host.parse::<IpAddr>().is_err() {
builder = builder.resolve_to_addrs(host, addrs);
}
builder
.build()
.map_err(|e| ToolError::ExecutionFailed(format!("failed to build HTTP client: {}", e)))
}
const MAX_REDIRECTS: usize = 10;
pub async fn guarded_get(
url: &str,
check_ssrf: bool,
timeout: Option<Duration>,
) -> Result<reqwest::Response, ToolError> {
let mut current = url.to_string();
for _ in 0..=MAX_REDIRECTS {
let parsed = parse_http_url(¤t)?;
let host = parsed.host_str().expect("host checked above").to_string();
let addrs = resolve_url_addrs(&parsed).await?;
if check_ssrf {
ensure_all_public(&addrs)?;
}
let client = pinned_client(&host, &addrs, timeout)?;
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
)))
}
pub async fn guarded_post_json(
url: &str,
body: &serde_json::Value,
check_ssrf: bool,
timeout: Option<Duration>,
) -> Result<reqwest::Response, ToolError> {
let parsed = parse_http_url(url)?;
let host = parsed.host_str().expect("host checked above").to_string();
let addrs = resolve_url_addrs(&parsed).await?;
if check_ssrf {
ensure_all_public(&addrs)?;
}
let client = pinned_client(&host, &addrs, timeout)?;
client
.post(url)
.json(body)
.send()
.await
.map_err(|e| ToolError::ExecutionFailed(format!("HTTP request failed: {}", e)))
}
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()));
assert!(!is_private_ip(
&"2606:4700:4700::1111".parse::<IpAddr>().unwrap()
));
}
#[test]
fn ipv4_special_ranges_are_blocked() {
let blocked: &[&str] = &[
"100.64.0.1", "100.127.255.1", "198.18.0.1", "198.19.255.1", "192.0.0.1", "192.0.2.1", "198.51.100.1", "203.0.113.1", "224.0.0.1", "240.0.0.1", "0.1.2.3", ];
for s in blocked {
assert!(
is_private_ip(&s.parse::<IpAddr>().unwrap()),
"expected {s} to be flagged"
);
}
for s in &["100.128.0.1", "198.20.0.1", "8.8.8.8", "1.1.1.1"] {
assert!(
!is_private_ip(&s.parse::<IpAddr>().unwrap()),
"expected {s} to be allowed"
);
}
}
#[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());
}
#[test]
fn parse_http_url_rejects_scheme_and_missing_host() {
assert!(parse_http_url("file:///etc/passwd").is_err());
assert!(parse_http_url("ftp://b.com/x").is_err());
assert!(parse_http_url("not a url").is_err());
assert!(parse_http_url("http://").is_err());
assert!(parse_http_url("https://example.com/x").is_ok());
}
#[test]
fn ensure_all_public_blocks_when_any_answer_is_private() {
let mixed: Vec<SocketAddr> = vec![
"8.8.8.8:443".parse().unwrap(),
"10.0.0.5:443".parse().unwrap(),
"1.1.1.1:443".parse().unwrap(),
];
let err = ensure_all_public(&mixed).unwrap_err();
assert!(err.to_string().contains("SSRF"), "got: {err}");
let public: Vec<SocketAddr> = vec![
"8.8.8.8:443".parse().unwrap(),
"[2606:4700:4700::1111]:443".parse().unwrap(),
];
ensure_all_public(&public).expect("all-public set passes");
let mapped: Vec<SocketAddr> = vec!["[::ffff:169.254.169.254]:80".parse().unwrap()];
assert!(ensure_all_public(&mapped).is_err());
}
#[tokio::test]
async fn resolve_url_addrs_ip_literals_bypass_dns() {
let url = parse_http_url("http://127.0.0.1:8080/").unwrap();
let addrs = resolve_url_addrs(&url).await.unwrap();
assert_eq!(addrs, vec!["127.0.0.1:8080".parse::<SocketAddr>().unwrap()]);
let url = parse_http_url("https://8.8.8.8/").unwrap();
let addrs = resolve_url_addrs(&url).await.unwrap();
assert_eq!(addrs, vec!["8.8.8.8:443".parse::<SocketAddr>().unwrap()]);
}
#[tokio::test]
async fn guarded_get_blocks_loopback_before_connecting() {
let err = guarded_get("http://127.0.0.1:1/", true, None)
.await
.unwrap_err();
assert!(err.to_string().contains("SSRF"), "got: {err}");
}
#[tokio::test]
async fn guarded_get_blocks_link_local_before_connecting() {
let err = guarded_get("http://[::ffff:169.254.169.254]/latest", true, None)
.await
.unwrap_err();
assert!(err.to_string().contains("SSRF"), "got: {err}");
}
#[tokio::test]
async fn guarded_get_rejects_non_http_scheme_without_sending() {
let err = guarded_get("file:///etc/passwd", true, None)
.await
.unwrap_err();
assert!(err.to_string().contains("scheme"), "got: {err}");
}
#[tokio::test]
async fn guarded_post_json_blocks_private_before_connecting() {
let err = guarded_post_json(
"http://169.254.169.254/latest/meta-data/",
&serde_json::json!({}),
true,
None,
)
.await
.unwrap_err();
assert!(err.to_string().contains("SSRF"), "got: {err}");
}
}