use std::net::{Ipv4Addr, Ipv6Addr};
use chat_engine_sdk::error::PluginError;
use url::{Host, Url};
pub fn validate_outbound_url(raw: &str, key_name: &str) -> Result<Url, PluginError> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Err(PluginError::invalid_input(format!(
"{key_name} must not be empty",
)));
}
let url = Url::parse(trimmed).map_err(|e| {
PluginError::invalid_input_with(format!("{key_name} is not a valid absolute URL"), e)
})?;
if url.scheme() != "https" {
return Err(PluginError::invalid_input(format!(
"{key_name} must use the https scheme; got `{}`",
url.scheme(),
)));
}
let host = url.host().ok_or_else(|| {
PluginError::invalid_input(format!("{key_name} must include a host component"))
})?;
match host {
Host::Ipv4(v4) => {
if is_disallowed_ipv4(v4) {
return Err(PluginError::invalid_input(format!(
"{key_name} resolves to a disallowed IPv4 address range \
(loopback / link-local / private / multicast / reserved)",
)));
}
}
Host::Ipv6(v6) => {
if is_disallowed_ipv6(v6) {
return Err(PluginError::invalid_input(format!(
"{key_name} resolves to a disallowed IPv6 address range \
(loopback / link-local / unique-local / multicast / reserved)",
)));
}
}
Host::Domain(name) => {
let lower = name.trim_end_matches('.').to_ascii_lowercase();
if lower == "localhost" || lower.ends_with(".localhost") {
return Err(PluginError::invalid_input(format!(
"{key_name} must not point at the loopback hostname (`localhost`)",
)));
}
}
}
Ok(url)
}
fn is_disallowed_ipv4(addr: Ipv4Addr) -> bool {
let o = addr.octets();
addr.is_loopback()
|| addr.is_private()
|| addr.is_unspecified()
|| addr.is_multicast()
|| addr.is_broadcast()
|| addr.is_documentation()
|| o[0] == 0
|| (o[0] == 169 && o[1] == 254)
|| (o[0] == 100 && (o[1] & 0xc0) == 64)
}
fn is_disallowed_ipv6(addr: Ipv6Addr) -> bool {
if addr.is_loopback() || addr.is_unspecified() || addr.is_multicast() {
return true;
}
let segs = addr.segments();
if (segs[0] & 0xfe00) == 0xfc00 {
return true;
}
if (segs[0] & 0xffc0) == 0xfe80 {
return true;
}
if let Some(v4) = addr.to_ipv4_mapped() {
return is_disallowed_ipv4(v4);
}
if let Some(v4) = addr.to_ipv4() {
return is_disallowed_ipv4(v4);
}
false
}
#[cfg(test)]
#[path = "url_guard_tests.rs"]
mod url_guard_tests;