use crate::error::A2aError;
use std::net::{Ipv4Addr, Ipv6Addr};
use url::{Host, Url};
const REQUIRED_SCHEME: &str = "https";
const IPV6_UNIQUE_LOCAL_MASK: u16 = 0xfe00;
const IPV6_UNIQUE_LOCAL_PREFIX: u16 = 0xfc00;
const IPV6_LINK_LOCAL_MASK: u16 = 0xffc0;
const IPV6_LINK_LOCAL_PREFIX: u16 = 0xfe80;
const DISALLOWED_HOSTNAMES: [&str; 2] = ["localhost", "localhost.localdomain"];
pub fn validate_callback_url(raw: &str) -> Result<(), A2aError> {
let parsed = Url::parse(raw).map_err(|e| {
A2aError::InvalidParams(format!("invalid push-notification callback URL: {e}"))
})?;
if parsed.scheme() != REQUIRED_SCHEME {
return Err(A2aError::InvalidParams(format!(
"push-notification callback URL must use {REQUIRED_SCHEME}; got scheme `{}`",
parsed.scheme()
)));
}
let host = parsed.host().ok_or_else(|| {
A2aError::InvalidParams("push-notification callback URL has no host".into())
})?;
if is_disallowed_host(&host) {
return Err(A2aError::InvalidParams(format!(
"push-notification callback URL host `{}` is not allowed \
(loopback/private/link-local/reserved ranges are blocked)",
parsed.host_str().unwrap_or_default(),
)));
}
Ok(())
}
fn is_disallowed_host(host: &Host<&str>) -> bool {
match host {
Host::Domain(name) => DISALLOWED_HOSTNAMES.contains(&name.to_ascii_lowercase().as_str()),
Host::Ipv4(v4) => is_disallowed_ipv4(v4),
Host::Ipv6(v6) => is_disallowed_ipv6(v6),
}
}
fn is_disallowed_ipv4(v4: &Ipv4Addr) -> bool {
v4.is_loopback()
|| v4.is_private()
|| v4.is_link_local()
|| v4.is_unspecified()
|| v4.is_broadcast()
|| v4.is_documentation()
|| v4.is_multicast()
}
fn is_disallowed_ipv6(v6: &Ipv6Addr) -> bool {
if v6.is_loopback() || v6.is_unspecified() || v6.is_multicast() {
return true;
}
if let Some(mapped) = v6.to_ipv4_mapped() {
return is_disallowed_ipv4(&mapped);
}
let first_segment = v6.segments()[0];
let is_unique_local = (first_segment & IPV6_UNIQUE_LOCAL_MASK) == IPV6_UNIQUE_LOCAL_PREFIX;
let is_link_local = (first_segment & IPV6_LINK_LOCAL_MASK) == IPV6_LINK_LOCAL_PREFIX;
is_unique_local || is_link_local
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn https_public_host_is_allowed() {
validate_callback_url("https://example.test/hook").expect("public https URL must pass");
}
#[test]
fn plain_http_is_rejected() {
let err = validate_callback_url("http://example.test/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn loopback_hostname_is_rejected() {
let err = validate_callback_url("https://localhost/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn loopback_ipv4_literal_is_rejected() {
let err = validate_callback_url("https://127.0.0.1/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn loopback_ipv6_literal_is_rejected() {
let err = validate_callback_url("https://[::1]/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn private_ipv4_ranges_are_rejected() {
for host in ["10.0.0.1", "172.16.0.1", "192.168.1.1"] {
let url = format!("https://{host}/hook");
let err = validate_callback_url(&url).unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)), "{host}");
}
}
#[test]
fn link_local_metadata_endpoint_is_rejected() {
let err = validate_callback_url("https://169.254.169.254/latest/meta-data/").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn ipv4_mapped_ipv6_loopback_is_rejected() {
let err = validate_callback_url("https://[::ffff:127.0.0.1]/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn ipv6_unique_local_is_rejected() {
let err = validate_callback_url("https://[fc00::1]/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn ipv6_link_local_is_rejected() {
let err = validate_callback_url("https://[fe80::1]/hook").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn public_ipv4_literal_is_allowed() {
validate_callback_url("https://93.184.216.34/hook").expect("public IP must pass");
}
#[test]
fn malformed_url_is_rejected() {
let err = validate_callback_url("not a url").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
#[test]
fn url_with_empty_host_is_rejected() {
let err = validate_callback_url("https://").unwrap_err();
assert!(matches!(err, A2aError::InvalidParams(_)));
}
}