Skip to main content

mmdb_writer/
net.rs

1//! IP-version handling and the conversion from [`ipnet::IpNet`] to the `(bits, prefix_len)`
2//! pair the tree walks.
3
4// Narrowing casts are intentional: addresses and prefix lengths are range-bounded by the IP
5// family before the cast (a v4 base is < 2^32, a prefix length <= 128).
6#![allow(clippy::cast_possible_truncation)]
7
8use std::net::IpAddr;
9
10use ipnet::{IpNet, Ipv4Net, Ipv6Net};
11
12use crate::error::Error;
13
14/// Which IP version a database indexes.
15///
16/// A [`Writer`](crate::Writer) defaults to [`IpVersion::V6`], which stores IPv4 networks
17/// inside the IPv4-in-IPv6 range (`::/96`) so a single database answers both IPv4 and IPv6
18/// lookups. Choose [`IpVersion::V4`] for a smaller, IPv4-only database that rejects IPv6
19/// inserts.
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
21#[non_exhaustive]
22pub enum IpVersion {
23    /// A 32-bit tree. Only IPv4 networks may be inserted.
24    V4,
25    /// A 128-bit tree. IPv4 networks are stored under `::/96`. **Default.**
26    #[default]
27    V6,
28}
29
30impl IpVersion {
31    /// Depth of the search tree in bits (32 for IPv4, 128 for IPv6).
32    pub(crate) const fn tree_depth(self) -> u8 {
33        match self {
34            Self::V4 => 32,
35            Self::V6 => 128,
36        }
37    }
38
39    /// The value written to the metadata `ip_version` field (4 or 6).
40    pub(crate) const fn metadata(self) -> u16 {
41        match self {
42            Self::V4 => 4,
43            Self::V6 => 6,
44        }
45    }
46}
47
48/// Convert a network to the `(bits, prefix_len)` pair used to walk the tree.
49///
50/// Host bits are truncated first (matching the Go writer). In a V6 tree, IPv4 networks are
51/// lifted into `::/96`. In a V4 tree, the address occupies the top 32 bits of the walk word
52/// and IPv6 inputs are rejected.
53pub(crate) fn to_tree_prefix(net: IpNet, ip_version: IpVersion) -> Result<(u128, u8), Error> {
54    match (ip_version, net.trunc()) {
55        (IpVersion::V6, IpNet::V6(v6)) => {
56            let bits = u128::from_be_bytes(v6.addr().octets());
57            Ok((bits, v6.prefix_len()))
58        }
59        (IpVersion::V6, IpNet::V4(v4)) => {
60            // IPv4 lives under `::/96`: leave the top 96 bits zero.
61            let v4_bits = u32::from_be_bytes(v4.addr().octets());
62            Ok((u128::from(v4_bits), 96 + v4.prefix_len()))
63        }
64        (IpVersion::V4, IpNet::V4(v4)) => {
65            // The address occupies the top 32 bits so `bits >> (127 - depth)` reads it
66            // most-significant-bit first over depths 0..32.
67            let v4_bits = u32::from_be_bytes(v4.addr().octets());
68            Ok((u128::from(v4_bits) << 96, v4.prefix_len()))
69        }
70        (IpVersion::V4, IpNet::V6(v6)) => Err(Error::Ipv6InIpv4Tree(v6)),
71    }
72}
73
74/// The three IPv6 networks that are aliased to the IPv4 subtree: IPv4-mapped, Teredo, and
75/// 6to4. Used to reject inserts that would land in aliased space.
76pub(crate) fn alias_networks() -> [IpNet; 3] {
77    [
78        "::ffff:0:0/96".parse().expect("valid alias CIDR"),
79        "2001::/32".parse().expect("valid alias CIDR"),
80        "2002::/16".parse().expect("valid alias CIDR"),
81    ]
82}
83
84/// Decompose an inclusive address range `[start, end]` into the minimal list of CIDR
85/// networks that exactly covers it.
86///
87/// Both endpoints must be the same IP family. The result is ordered from `start` upward.
88pub(crate) fn range_to_networks(start: IpAddr, end: IpAddr) -> Result<Vec<IpNet>, Error> {
89    match (start, end) {
90        (IpAddr::V4(s), IpAddr::V4(e)) => {
91            let (s, e) = (u32::from(s), u32::from(e));
92            if s > e {
93                return Err(Error::InvalidRange(
94                    "start address is greater than end address",
95                ));
96            }
97            Ok(range_to_cidrs(u128::from(s), u128::from(e), 32)
98                .into_iter()
99                .map(|(base, prefix)| {
100                    // `base` fits in 32 bits and `prefix <= 32`, so both casts are exact.
101                    let addr = std::net::Ipv4Addr::from(base as u32);
102                    IpNet::V4(Ipv4Net::new(addr, prefix).expect("prefix within range"))
103                })
104                .collect())
105        }
106        (IpAddr::V6(s), IpAddr::V6(e)) => {
107            let (s, e) = (u128::from(s), u128::from(e));
108            if s > e {
109                return Err(Error::InvalidRange(
110                    "start address is greater than end address",
111                ));
112            }
113            Ok(range_to_cidrs(s, e, 128)
114                .into_iter()
115                .map(|(base, prefix)| {
116                    let addr = std::net::Ipv6Addr::from(base);
117                    IpNet::V6(Ipv6Net::new(addr, prefix).expect("prefix within range"))
118                })
119                .collect())
120        }
121        _ => Err(Error::InvalidRange(
122            "start and end are different IP families",
123        )),
124    }
125}
126
127/// Core integer range → CIDR decomposition, operating in a `total_bits`-wide address space
128/// (32 for IPv4, 128 for IPv6). Returns `(base, prefix_len)` pairs.
129fn range_to_cidrs(mut start: u128, end: u128, total_bits: u32) -> Vec<(u128, u8)> {
130    // Whole-space shortcut avoids a `1 << total_bits` overflow below.
131    if start == 0 && end == max_for_bits(total_bits) {
132        return vec![(0, 0)];
133    }
134    let mut cidrs = Vec::new();
135    loop {
136        // Largest block whose size is limited by `start`'s alignment...
137        let align_bits = if start == 0 {
138            total_bits
139        } else {
140            start.trailing_zeros().min(total_bits)
141        };
142        // ...and shrunk until the block fits within `end`.
143        let mut size = align_bits;
144        while size > 0 {
145            let block = 1u128 << size;
146            let block_end = start + (block - 1);
147            if block_end <= end {
148                break;
149            }
150            size -= 1;
151        }
152        let prefix_len = (total_bits - size) as u8;
153        cidrs.push((start, prefix_len));
154
155        let block = 1u128 << size;
156        match start.checked_add(block) {
157            Some(next) if next <= end => start = next,
158            _ => break,
159        }
160    }
161    cidrs
162}
163
164fn max_for_bits(total_bits: u32) -> u128 {
165    if total_bits >= 128 {
166        u128::MAX
167    } else {
168        (1u128 << total_bits) - 1
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use super::*;
175
176    fn v4(s: &str) -> IpAddr {
177        s.parse().unwrap()
178    }
179
180    #[test]
181    fn aligned_range_is_single_cidr() {
182        let nets = range_to_networks(v4("10.0.0.0"), v4("10.0.0.255")).unwrap();
183        assert_eq!(nets, vec!["10.0.0.0/24".parse().unwrap()]);
184    }
185
186    #[test]
187    fn unaligned_range_splits() {
188        // 10.0.0.1 - 10.0.0.2 → two /32s.
189        let nets = range_to_networks(v4("10.0.0.1"), v4("10.0.0.2")).unwrap();
190        assert_eq!(
191            nets,
192            vec![
193                "10.0.0.1/32".parse().unwrap(),
194                "10.0.0.2/32".parse().unwrap()
195            ]
196        );
197    }
198
199    #[test]
200    fn classic_range_decomposition() {
201        // 1.1.1.0 - 1.1.1.5 → /29? no: 0..5 = 0-3 (/30) + 4-5 (/31).
202        let nets = range_to_networks(v4("1.1.1.0"), v4("1.1.1.5")).unwrap();
203        assert_eq!(
204            nets,
205            vec!["1.1.1.0/30".parse().unwrap(), "1.1.1.4/31".parse().unwrap()]
206        );
207    }
208
209    #[test]
210    fn single_address_range() {
211        let nets = range_to_networks(v4("192.168.1.1"), v4("192.168.1.1")).unwrap();
212        assert_eq!(nets, vec!["192.168.1.1/32".parse().unwrap()]);
213    }
214
215    #[test]
216    fn whole_ipv4_space() {
217        let nets = range_to_networks(v4("0.0.0.0"), v4("255.255.255.255")).unwrap();
218        assert_eq!(nets, vec!["0.0.0.0/0".parse().unwrap()]);
219    }
220
221    #[test]
222    fn reversed_range_errors() {
223        assert!(range_to_networks(v4("10.0.0.5"), v4("10.0.0.1")).is_err());
224    }
225
226    #[test]
227    fn mismatched_families_error() {
228        assert!(range_to_networks(v4("10.0.0.0"), "::1".parse().unwrap()).is_err());
229    }
230
231    #[test]
232    fn ipv6_range() {
233        let nets = range_to_networks(
234            "2001:db8::".parse().unwrap(),
235            "2001:db8::ff".parse().unwrap(),
236        )
237        .unwrap();
238        assert_eq!(nets, vec!["2001:db8::/120".parse().unwrap()]);
239    }
240
241    #[test]
242    fn ipv6_single_address_range() {
243        let ip: IpAddr = "2001:db8::7".parse().unwrap();
244        assert_eq!(
245            range_to_networks(ip, ip).unwrap(),
246            vec!["2001:db8::7/128".parse().unwrap()]
247        );
248    }
249
250    #[test]
251    fn ipv6_reversed_range_errors() {
252        assert!(
253            range_to_networks(
254                "2001:db8::2".parse().unwrap(),
255                "2001:db8::1".parse().unwrap()
256            )
257            .is_err()
258        );
259    }
260
261    #[test]
262    fn whole_ipv6_space() {
263        let nets = range_to_networks(
264            "::".parse().unwrap(),
265            "ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff".parse().unwrap(),
266        )
267        .unwrap();
268        assert_eq!(nets, vec!["::/0".parse().unwrap()]);
269    }
270
271    #[test]
272    fn zero_start_partial_range_is_not_whole_space() {
273        // Starts at 0.0.0.0 but does not reach the top: must NOT shortcut to /0.
274        let nets = range_to_networks(v4("0.0.0.0"), v4("0.0.0.7")).unwrap();
275        assert_eq!(nets, vec!["0.0.0.0/29".parse().unwrap()]);
276        // Ends at the top but does not start at 0: also not /0.
277        let nets = range_to_networks(v4("255.255.255.248"), v4("255.255.255.255")).unwrap();
278        assert_eq!(nets, vec!["255.255.255.248/29".parse().unwrap()]);
279    }
280}