use std::net::Ipv4Addr;
mod generated;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub enum CloudProvider {
Aws,
Gcp,
Cloudflare,
}
impl CloudProvider {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Aws => "aws",
Self::Gcp => "gcp",
Self::Cloudflare => "cloudflare",
}
}
}
#[must_use]
pub fn classify_ipv4(ip: Ipv4Addr) -> Option<CloudProvider> {
let key = u32::from(ip);
let ranges = generated::CLOUD_RANGES;
let idx = ranges.partition_point(|(start, _, _)| *start <= key);
if idx == 0 {
return None;
}
let (start, end, provider) = ranges[idx - 1];
if key >= start && key <= end {
Some(provider)
} else {
None
}
}
#[cfg(test)]
mod tests {
#[test]
fn classify_ipv4_below_first_range_is_none() {
use std::net::Ipv4Addr;
assert_eq!(classify_ipv4(Ipv4Addr::UNSPECIFIED), None);
}
use super::*;
#[test]
fn known_aws_address_classifies_as_aws() {
assert_eq!(
classify_ipv4(Ipv4Addr::new(1, 178, 1, 0)),
Some(CloudProvider::Aws)
);
}
#[test]
fn known_gcp_address_classifies_as_gcp() {
assert_eq!(
classify_ipv4(Ipv4Addr::new(8, 34, 208, 0)),
Some(CloudProvider::Gcp)
);
}
#[test]
fn known_cloudflare_address_classifies_as_cloudflare() {
assert_eq!(
classify_ipv4(Ipv4Addr::new(103, 21, 244, 0)),
Some(CloudProvider::Cloudflare)
);
}
#[test]
fn private_address_is_unknown() {
assert_eq!(classify_ipv4(Ipv4Addr::new(10, 0, 0, 1)), None);
}
#[test]
fn ranges_are_disjoint_and_sorted() {
let r = generated::CLOUD_RANGES;
for w in r.windows(2) {
assert!(w[0].0 <= w[0].1, "interval start must be <= end");
assert!(w[0].1 < w[1].0, "intervals must be disjoint and ascending");
}
}
#[test]
fn as_str_tokens() {
assert_eq!(CloudProvider::Aws.as_str(), "aws");
assert_eq!(CloudProvider::Gcp.as_str(), "gcp");
assert_eq!(CloudProvider::Cloudflare.as_str(), "cloudflare");
}
}