1#![allow(clippy::cast_possible_truncation)]
7
8use std::net::IpAddr;
9
10use ipnet::{IpNet, Ipv4Net, Ipv6Net};
11
12use crate::error::Error;
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
21#[non_exhaustive]
22pub enum IpVersion {
23 V4,
25 #[default]
27 V6,
28}
29
30impl IpVersion {
31 pub(crate) const fn tree_depth(self) -> u8 {
33 match self {
34 Self::V4 => 32,
35 Self::V6 => 128,
36 }
37 }
38
39 pub(crate) const fn metadata(self) -> u16 {
41 match self {
42 Self::V4 => 4,
43 Self::V6 => 6,
44 }
45 }
46}
47
48pub(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 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 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
74pub(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
84pub(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 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
127fn range_to_cidrs(mut start: u128, end: u128, total_bits: u32) -> Vec<(u128, u8)> {
130 if start == 0 && end == max_for_bits(total_bits) {
132 return vec![(0, 0)];
133 }
134 let mut cidrs = Vec::new();
135 loop {
136 let align_bits = if start == 0 {
138 total_bits
139 } else {
140 start.trailing_zeros().min(total_bits)
141 };
142 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 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 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 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 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}