use std::io;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, TcpStream, ToSocketAddrs};
use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SsrfError {
pub host: String,
pub reason: String,
}
impl std::fmt::Display for SsrfError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"host `{}` rejected by SSRF guard: {}",
self.host, self.reason
)
}
}
impl std::error::Error for SsrfError {}
fn reject(host: &str, reason: impl Into<String>) -> SsrfError {
SsrfError {
host: host.to_string(),
reason: reason.into(),
}
}
pub fn is_global(ip: IpAddr) -> bool {
match ip {
IpAddr::V4(v4) => is_global_v4(v4),
IpAddr::V6(v6) => is_global_v6(v6),
}
}
fn is_global_v4(ip: Ipv4Addr) -> bool {
let [a, b, _, _] = ip.octets();
if a == 0 {
return false;
}
if ip.is_loopback()
|| ip.is_private()
|| ip.is_link_local()
|| ip.is_broadcast()
|| ip.is_multicast()
|| ip.is_unspecified()
|| ip.is_documentation()
{
return false;
}
if a == 100 && (64..=127).contains(&b) {
return false;
}
if a >= 240 {
return false;
}
true
}
fn is_global_v6(ip: Ipv6Addr) -> bool {
if let Some(v4) = ip.to_ipv4_mapped() {
return is_global_v4(v4);
}
if let Some(v4) = ip.to_ipv4() {
return is_global_v4(v4);
}
if ip.is_loopback() || ip.is_unspecified() || ip.is_multicast() {
return false;
}
let segments = ip.segments();
if (segments[0] & 0xffc0) == 0xfe80 {
return false;
}
if (segments[0] & 0xfe00) == 0xfc00 {
return false;
}
if segments[0] == 0x2001 && segments[1] == 0x0db8 {
return false;
}
true
}
pub fn guard_host(host: &str, allow_private: bool) -> Result<(), SsrfError> {
if allow_private {
return Ok(());
}
resolve_guarded(host, 0, false).map(|_| ())
}
pub type ResolveFn = fn(&str, u16) -> io::Result<Vec<SocketAddr>>;
pub fn std_resolve(host: &str, port: u16) -> io::Result<Vec<SocketAddr>> {
let bare = host
.strip_prefix('[')
.and_then(|s| s.strip_suffix(']'))
.unwrap_or(host);
if let Ok(ip) = bare.parse::<IpAddr>() {
return Ok(vec![SocketAddr::new(ip, port)]);
}
(host, port).to_socket_addrs().map(|it| it.collect())
}
pub fn resolve_guarded(
host: &str,
port: u16,
allow_private: bool,
) -> Result<Vec<SocketAddr>, SsrfError> {
resolve_guarded_with(host, port, allow_private, std_resolve)
}
pub fn resolve_guarded_with(
host: &str,
port: u16,
allow_private: bool,
resolve: ResolveFn,
) -> Result<Vec<SocketAddr>, SsrfError> {
if host.is_empty() {
return Err(reject(host, "empty host"));
}
let addrs = resolve(host, port).map_err(|e| reject(host, format!("resolve failed: {e}")))?;
if addrs.is_empty() {
return Err(reject(host, "no addresses resolved"));
}
if !allow_private {
for sa in &addrs {
check_addr(host, sa.ip())?;
}
}
Ok(addrs)
}
pub fn connect_addrs(
host: &str,
addrs: &[SocketAddr],
timeout: Duration,
allow_private: bool,
) -> io::Result<TcpStream> {
if addrs.is_empty() {
return Err(io::Error::new(
io::ErrorKind::NotFound,
format!("no vetted addresses for {host}"),
));
}
if !allow_private {
for sa in addrs {
if let Err(e) = check_addr(host, sa.ip()) {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
e.to_string(),
));
}
}
}
let mut last: Option<io::Error> = None;
for sa in addrs {
match TcpStream::connect_timeout(sa, timeout) {
Ok(stream) => {
stream.set_read_timeout(Some(timeout))?;
stream.set_write_timeout(Some(timeout))?;
stream.set_nodelay(true).ok();
return Ok(stream);
}
Err(e) => last = Some(e),
}
}
Err(last.unwrap_or_else(|| {
io::Error::new(io::ErrorKind::NotFound, format!("cannot connect to {host}"))
}))
}
pub fn connect_vetted(
host: &str,
port: u16,
timeout: Duration,
allow_private: bool,
) -> io::Result<TcpStream> {
connect_vetted_with(host, port, timeout, allow_private, std_resolve)
}
pub fn connect_vetted_with(
host: &str,
port: u16,
timeout: Duration,
allow_private: bool,
resolve: ResolveFn,
) -> io::Result<TcpStream> {
let addrs = resolve_guarded_with(host, port, allow_private, resolve)
.map_err(|e| io::Error::new(io::ErrorKind::PermissionDenied, e.to_string()))?;
connect_addrs(host, &addrs, timeout, allow_private)
}
fn check_addr(host: &str, ip: IpAddr) -> Result<(), SsrfError> {
if is_global(ip) {
Ok(())
} else {
Err(reject(host, format!("{} ({ip})", class_of(ip))))
}
}
fn class_of(ip: IpAddr) -> &'static str {
match ip {
IpAddr::V4(v4) => {
if v4.is_unspecified() {
"unspecified"
} else if v4.is_loopback() {
"loopback"
} else if v4.is_private() {
"private (RFC-1918)"
} else if v4.is_link_local() {
"link-local"
} else if v4.is_broadcast() {
"broadcast"
} else if v4.is_multicast() {
"multicast"
} else {
"reserved"
}
}
IpAddr::V6(v6) => {
if let Some(v4) = v6.to_ipv4_mapped().or_else(|| v6.to_ipv4()) {
return class_of(IpAddr::V4(v4));
}
if v6.is_unspecified() {
"unspecified"
} else if v6.is_loopback() {
"loopback"
} else if v6.is_multicast() {
"multicast"
} else {
"link-local/unique-local"
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn v4(a: u8, b: u8, c: u8, d: u8) -> IpAddr {
IpAddr::V4(Ipv4Addr::new(a, b, c, d))
}
fn v6(s: &str) -> IpAddr {
IpAddr::V6(s.parse::<Ipv6Addr>().expect("test ipv6 literal"))
}
#[test]
fn public_v4_is_global() {
assert!(is_global(v4(8, 8, 8, 8)));
assert!(is_global(v4(1, 1, 1, 1)));
assert!(is_global(v4(93, 184, 216, 34))); assert!(is_global(v4(172, 15, 255, 255))); assert!(is_global(v4(172, 32, 0, 1))); assert!(is_global(v4(11, 0, 0, 1))); assert!(is_global(v4(192, 167, 255, 255))); assert!(is_global(v4(192, 169, 0, 1))); assert!(is_global(v4(100, 63, 255, 255))); assert!(is_global(v4(100, 128, 0, 1))); }
#[test]
fn loopback_blocked() {
assert!(!is_global(v4(127, 0, 0, 1)));
assert!(!is_global(v4(127, 255, 255, 255)));
assert!(!is_global(v6("::1")));
}
#[test]
fn rfc1918_blocked() {
assert!(!is_global(v4(10, 0, 0, 0)));
assert!(!is_global(v4(10, 255, 255, 255)));
assert!(!is_global(v4(172, 16, 0, 0)));
assert!(!is_global(v4(172, 16, 0, 1)));
assert!(!is_global(v4(172, 31, 255, 255)));
assert!(!is_global(v4(192, 168, 0, 1)));
assert!(!is_global(v4(192, 168, 255, 255)));
}
#[test]
fn link_local_and_metadata_blocked() {
assert!(!is_global(v4(169, 254, 0, 1)));
assert!(!is_global(v4(169, 254, 169, 254)));
assert!(!is_global(v4(169, 254, 255, 255)));
assert!(!is_global(v6("fe80::1")));
assert!(!is_global(v6("febf:ffff:ffff:ffff:ffff:ffff:ffff:ffff")));
}
#[test]
fn unique_local_blocked() {
assert!(!is_global(v6("fc00::1")));
assert!(!is_global(v6("fd12:3456:789a::1")));
assert!(!is_global(v6("fdff:ffff:ffff:ffff:ffff:ffff:ffff:ffff")));
}
#[test]
fn unspecified_blocked() {
assert!(!is_global(v4(0, 0, 0, 0)));
assert!(!is_global(v4(0, 1, 2, 3))); assert!(!is_global(v6("::")));
}
#[test]
fn multicast_and_broadcast_blocked() {
assert!(!is_global(v4(224, 0, 0, 1)));
assert!(!is_global(v4(239, 255, 255, 255)));
assert!(!is_global(v4(255, 255, 255, 255))); assert!(!is_global(v6("ff02::1")));
}
#[test]
fn reserved_v4_blocked() {
assert!(!is_global(v4(240, 0, 0, 1)));
assert!(!is_global(v4(255, 0, 0, 1)));
}
#[test]
fn ipv4_mapped_bypass_is_caught() {
assert!(!is_global(v6("::ffff:127.0.0.1")));
assert!(!is_global(v6("::ffff:10.0.0.1")));
assert!(!is_global(v6("::ffff:169.254.169.254")));
assert!(!is_global(v6("::ffff:192.168.1.1")));
assert!(is_global(v6("::ffff:8.8.8.8")));
}
#[test]
fn ipv4_compatible_bypass_is_caught() {
assert!(!is_global(v6("::10.0.0.1")));
assert!(!is_global(v6("::169.254.169.254")));
}
#[test]
fn public_v6_is_global() {
assert!(is_global(v6("2606:4700:4700::1111"))); assert!(is_global(v6("2001:4860:4860::8888"))); }
#[test]
fn ipv6_documentation_blocked() {
assert!(!is_global(v6("2001:db8::1")));
}
#[test]
fn guard_rejects_ip_literals() {
assert!(guard_host("127.0.0.1", false).is_err());
assert!(guard_host("10.0.0.5", false).is_err());
assert!(guard_host("169.254.169.254", false).is_err());
assert!(guard_host("::1", false).is_err());
assert!(guard_host("[::1]", false).is_err()); assert!(guard_host("[fe80::1]", false).is_err());
assert!(guard_host("::ffff:127.0.0.1", false).is_err());
}
#[test]
fn guard_allows_public_ip_literals() {
assert!(guard_host("8.8.8.8", false).is_ok());
assert!(guard_host("1.1.1.1", false).is_ok());
assert!(guard_host("[2606:4700:4700::1111]", false).is_ok());
}
#[test]
fn allow_private_skips_everything() {
assert!(guard_host("127.0.0.1", true).is_ok());
assert!(guard_host("10.0.0.5", true).is_ok());
assert!(guard_host("", true).is_ok());
assert!(guard_host("anything.invalid", true).is_ok());
}
#[test]
fn empty_host_rejected_when_guarded() {
assert!(guard_host("", false).is_err());
}
#[test]
fn error_carries_host_and_class() {
let err = guard_host("169.254.169.254", false).unwrap_err();
assert_eq!(err.host, "169.254.169.254");
assert!(err.reason.contains("link-local"), "reason: {}", err.reason);
let shown = err.to_string();
assert!(shown.contains("169.254.169.254"));
assert!(shown.contains("SSRF guard"));
}
#[test]
fn class_labels_are_specific() {
assert_eq!(class_of(v4(127, 0, 0, 1)), "loopback");
assert_eq!(class_of(v4(10, 0, 0, 1)), "private (RFC-1918)");
assert_eq!(class_of(v4(169, 254, 1, 1)), "link-local");
assert_eq!(class_of(v4(0, 0, 0, 0)), "unspecified");
assert_eq!(class_of(v4(224, 0, 0, 1)), "multicast");
assert_eq!(class_of(v4(255, 255, 255, 255)), "broadcast");
assert_eq!(class_of(v4(240, 0, 0, 1)), "reserved");
assert_eq!(class_of(v6("::ffff:10.0.0.1")), "private (RFC-1918)");
}
}