zond_engine/core/models/ip/
set.rs1use super::range::{IpError, IpRange, Ipv4Range, Ipv6Range};
24use std::{
25 net::{IpAddr, Ipv4Addr, Ipv6Addr},
26 str::FromStr,
27};
28
29#[derive(Debug, thiserror::Error)]
31pub enum IpSetError {
32 #[error("Invalid target in set: {0}")]
34 InvalidTarget(#[from] IpError),
35}
36
37#[derive(Debug, Clone, Default, PartialEq, Eq)]
45pub struct IpSet {
46 v4: Vec<Ipv4Range>,
47 v6: Vec<Ipv6Range>,
48 v4_dirty: bool,
49 v6_dirty: bool,
50}
51
52impl IpSet {
53 pub fn new() -> Self {
55 Self::default()
56 }
57
58 pub fn insert(&mut self, ip: IpAddr) {
64 match ip {
65 IpAddr::V4(v4) => self.push_v4_range(Ipv4Range::new(v4, v4).unwrap()),
66 IpAddr::V6(v6) => self.push_v6_range(Ipv6Range::new(v6, v6).unwrap()),
67 }
68 }
69
70 pub fn insert_range(&mut self, range: IpRange) {
72 match range {
73 IpRange::V4(r) => self.push_v4_range(r),
74 IpRange::V6(r) => self.push_v6_range(r),
75 }
76 }
77
78 pub fn push_v4_range(&mut self, range: Ipv4Range) {
80 self.v4.push(range);
81 self.v4_dirty = true;
82 }
83
84 pub fn push_v6_range(&mut self, range: Ipv6Range) {
86 self.v6.push(range);
87 self.v6_dirty = true;
88 }
89
90 pub fn canonicalize(&mut self) {
95 if self.v4_dirty {
96 if !self.v4.is_empty() {
97 self.merge_v4();
98 }
99 self.v4_dirty = false;
100 }
101 if self.v6_dirty {
102 if !self.v6.is_empty() {
103 self.merge_v6();
104 }
105 self.v6_dirty = false;
106 }
107 }
108
109 fn merge_v4(&mut self) {
110 self.v4.sort_by_key(|r| r.start_addr);
111 let mut merged: Vec<Ipv4Range> = Vec::with_capacity(self.v4.len());
112 let mut current = self.v4[0];
113
114 for next in self.v4.drain(1..) {
115 let curr_end = u32::from(current.end_addr);
116 let next_start = u32::from(next.start_addr);
117
118 if next_start <= curr_end.saturating_add(1) {
119 if next.end_addr > current.end_addr {
120 current.end_addr = next.end_addr;
121 }
122 } else {
123 merged.push(current);
124 current = next;
125 }
126 }
127 merged.push(current);
128 self.v4 = merged;
129 }
130
131 fn merge_v6(&mut self) {
132 self.v6.sort_by_key(|r| r.start_addr);
133 let mut merged: Vec<Ipv6Range> = Vec::with_capacity(self.v6.len());
134 let mut current = self.v6[0];
135
136 for next in self.v6.drain(1..) {
137 let curr_end = u128::from(current.end_addr);
138 let next_start = u128::from(next.start_addr);
139
140 if next_start <= curr_end.saturating_add(1) {
141 if next.end_addr > current.end_addr {
142 current.end_addr = next.end_addr;
143 }
144 } else {
145 merged.push(current);
146 current = next;
147 }
148 }
149 merged.push(current);
150 self.v6 = merged;
151 }
152
153 pub fn contains(&self, ip: &IpAddr) -> bool {
157 if !self.v4_dirty && !self.v6_dirty {
158 return self.contains_canonical(ip);
159 }
160 match ip {
161 IpAddr::V4(v4) => {
162 let target = u32::from(*v4);
163 self.v4.iter().any(|range| {
164 let start = u32::from(range.start_addr);
165 let end = u32::from(range.end_addr);
166 target >= start && target <= end
167 })
168 }
169 IpAddr::V6(v6) => {
170 let target = u128::from(*v6);
171 self.v6.iter().any(|range| {
172 let start = u128::from(range.start_addr);
173 let end = u128::from(range.end_addr);
174 target >= start && target <= end
175 })
176 }
177 }
178 }
179
180 pub fn len(&self) -> u128 {
182 if !self.v4_dirty && !self.v6_dirty {
183 self.len_canonical()
184 } else {
185 let mut temp = self.clone();
186 temp.canonicalize();
187 temp.len_canonical()
188 }
189 }
190
191 pub fn is_empty(&self) -> bool {
193 self.v4.is_empty() && self.v6.is_empty()
194 }
195
196 pub fn iter(&self) -> Box<dyn Iterator<Item = IpAddr> + Send + '_> {
198 if self.v4_dirty || self.v6_dirty {
199 let mut temp = self.clone();
200 temp.canonicalize();
201 temp.into_iter()
202 } else {
203 let v4_iter = self.v4.iter().flat_map(|range| range.to_iter());
204 let v6_iter = self.v6.iter().flat_map(|range| range.to_iter());
205 Box::new(v4_iter.chain(v6_iter))
206 }
207 }
208
209 pub fn contains_canonical(&self, ip: &IpAddr) -> bool {
217 debug_assert!(
218 !self.v4_dirty && !self.v6_dirty,
219 "IpSet must be canonicalized before calling contains_canonical"
220 );
221 match ip {
222 IpAddr::V4(v4) => {
223 let target = u32::from(*v4);
224 self.v4
225 .binary_search_by(|range| {
226 let start = u32::from(range.start_addr);
227 let end = u32::from(range.end_addr);
228 if target < start {
229 std::cmp::Ordering::Greater
230 } else if target > end {
231 std::cmp::Ordering::Less
232 } else {
233 std::cmp::Ordering::Equal
234 }
235 })
236 .is_ok()
237 }
238 IpAddr::V6(v6) => {
239 let target = u128::from(*v6);
240 self.v6
241 .binary_search_by(|range| {
242 let start = u128::from(range.start_addr);
243 let end = u128::from(range.end_addr);
244 if target < start {
245 std::cmp::Ordering::Greater
246 } else if target > end {
247 std::cmp::Ordering::Less
248 } else {
249 std::cmp::Ordering::Equal
250 }
251 })
252 .is_ok()
253 }
254 }
255 }
256
257 pub fn len_canonical(&self) -> u128 {
263 debug_assert!(
264 !self.v4_dirty && !self.v6_dirty,
265 "IpSet must be canonicalized before calling len_canonical"
266 );
267 let v4_len: u128 = self.v4.iter().map(|r| r.len() as u128).sum();
268 let v6_len: u128 = self.v6.iter().map(|r| r.len()).sum();
269 v4_len + v6_len
270 }
271
272 pub fn v4(&self) -> &[Ipv4Range] {
274 &self.v4
275 }
276
277 pub fn v6(&self) -> &[Ipv6Range] {
279 &self.v6
280 }
281}
282
283impl IntoIterator for IpSet {
288 type Item = IpAddr;
289 type IntoIter = Box<dyn Iterator<Item = IpAddr> + Send>;
290
291 fn into_iter(mut self) -> Self::IntoIter {
293 self.canonicalize();
294 let v4_iter = self.v4.into_iter().flat_map(|range| {
295 let start: u32 = range.start_addr.into();
296 let end: u32 = range.end_addr.into();
297 (start..=end).map(|ip| IpAddr::V4(Ipv4Addr::from(ip)))
298 });
299
300 let v6_iter = self.v6.into_iter().flat_map(|range| {
301 let start: u128 = range.start_addr.into();
302 let end: u128 = range.end_addr.into();
303 (start..=end).map(|ip| IpAddr::V6(Ipv6Addr::from(ip)))
304 });
305
306 Box::new(v4_iter.chain(v6_iter))
307 }
308}
309
310impl Extend<IpAddr> for IpSet {
311 fn extend<T: IntoIterator<Item = IpAddr>>(&mut self, iter: T) {
312 for ip in iter {
313 match ip {
314 IpAddr::V4(v4) => self.v4.push(Ipv4Range::new(v4, v4).unwrap()),
315 IpAddr::V6(v6) => self.v6.push(Ipv6Range::new(v6, v6).unwrap()),
316 }
317 }
318 self.v4_dirty = true;
319 self.v6_dirty = true;
320 }
321}
322
323impl FromIterator<IpAddr> for IpSet {
324 fn from_iter<I: IntoIterator<Item = IpAddr>>(iter: I) -> Self {
325 let mut set = IpSet::new();
326 set.extend(iter);
327 set.canonicalize();
328 set
329 }
330}
331
332impl FromIterator<IpRange> for IpSet {
333 fn from_iter<I: IntoIterator<Item = IpRange>>(iter: I) -> Self {
334 let mut set = IpSet::new();
335 for range in iter {
336 match range {
337 IpRange::V4(r) => set.v4.push(r),
338 IpRange::V6(r) => set.v6.push(r),
339 }
340 }
341 set.v4_dirty = true;
342 set.v6_dirty = true;
343 set.canonicalize();
344 set
345 }
346}
347
348impl FromIterator<IpSet> for IpSet {
349 fn from_iter<I: IntoIterator<Item = IpSet>>(iter: I) -> Self {
350 let mut master = IpSet::new();
351 for set in iter {
352 master.v4.extend(set.v4);
353 master.v6.extend(set.v6);
354 }
355 master.v4_dirty = true;
356 master.v6_dirty = true;
357 master.canonicalize();
358 master
359 }
360}
361
362impl From<IpAddr> for IpSet {
363 fn from(ip: IpAddr) -> Self {
364 let mut set = Self::new();
365 set.insert(ip);
366 set
367 }
368}
369
370impl From<IpRange> for IpSet {
371 fn from(range: IpRange) -> Self {
372 let mut set = Self::new();
373 set.insert_range(range);
374 set
375 }
376}
377
378impl TryFrom<&str> for IpSet {
379 type Error = IpSetError;
380 fn try_from(value: &str) -> Result<Self, Self::Error> {
381 let mut set = IpSet::new();
382 for part in value
383 .split([',', ' '])
384 .filter(|part| !part.trim().is_empty())
385 {
386 let range = part.parse::<IpRange>()?;
387 set.insert_range(range);
388 }
389 set.canonicalize();
390 Ok(set)
391 }
392}
393
394impl FromStr for IpSet {
395 type Err = IpSetError;
396 fn from_str(s: &str) -> Result<Self, Self::Err> {
397 Self::try_from(s)
398 }
399}
400
401#[cfg(test)]
411mod tests {
412 use super::*;
413
414 #[test]
415 fn lazy_merging_v4() {
416 let mut set = IpSet::new();
417 set.insert(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)));
418 set.insert(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 2)));
419
420 assert_eq!(set.v4.len(), 2);
422 assert!(set.v4_dirty);
423
424 set.canonicalize();
426 assert_eq!(set.len(), 2);
427 assert!(!set.v4_dirty);
428 assert_eq!(set.v4.len(), 1);
429 }
430
431 #[test]
432 fn set_battle_test_overlaps() {
433 let mut set = IpSet::new();
434 set.insert_range("10.0.0.10-10.0.0.20".parse().unwrap());
436 set.insert_range("10.0.0.5-10.0.0.15".parse().unwrap());
438 set.insert_range("10.0.0.15-10.0.0.25".parse().unwrap());
440 set.insert_range("10.0.0.30-10.0.0.40".parse().unwrap());
442 set.insert_range("10.0.0.0-10.0.0.50".parse().unwrap());
444
445 set.canonicalize();
446 assert_eq!(set.len(), 51);
447 assert_eq!(set.v4().len(), 1);
448 }
449
450 #[test]
451 fn set_ipv6_adjacency_boundary() {
452 let mut set = IpSet::new();
453 let max_v6 = Ipv6Addr::from(u128::MAX);
455 let max_minus_1 = Ipv6Addr::from(u128::MAX - 1);
456
457 set.insert(IpAddr::V6(max_minus_1));
458 set.insert(IpAddr::V6(max_v6));
459
460 set.canonicalize();
461 assert_eq!(set.len(), 2);
462 assert_eq!(set.v6().len(), 1);
463 }
464
465 #[test]
466 fn iteration_is_lazy_safe() {
467 let mut set = IpSet::new();
468 set.insert(IpAddr::V4(Ipv4Addr::from(1)));
469 set.insert(IpAddr::V4(Ipv4Addr::from(2)));
470
471 set.canonicalize();
472 let ips: Vec<IpAddr> = set.iter().collect();
473 assert_eq!(ips.len(), 2);
474 assert!(!set.v4_dirty);
475 }
476
477 #[test]
478 fn empty_set_canonical_is_fine() {
479 let mut set = IpSet::new();
480 set.canonicalize();
481 assert_eq!(set.len_canonical(), 0);
482 assert!(set.v4().is_empty());
483 }
484
485 #[test]
486 fn from_str_mixed_advanced() {
487 let set = IpSet::from_str("1.1.1.1/32, 1.1.1.1, ::1-::1, 10.0.0.1-10.0.0.2").unwrap();
488 assert_eq!(set.len(), 4);
490 }
491
492 #[test]
493 fn bulk_extend_efficiency() {
494 let mut set = IpSet::new();
495 let ips = (0..100).map(|i| IpAddr::V4(Ipv4Addr::from(i)));
496 set.extend(ips);
497
498 assert_eq!(set.v4.len(), 100);
499 set.canonicalize();
500 assert_eq!(set.v4.len(), 1);
501 assert_eq!(set.len(), 100);
502 }
503
504 #[test]
505 fn canonical_queries_panics_in_debug() {
506 #[cfg(debug_assertions)]
507 {
508 let set = IpSet::from_iter(vec![IpAddr::V4(Ipv4Addr::LOCALHOST)]);
509 assert!(!set.v4_dirty);
511 assert!(set.contains_canonical(&IpAddr::V4(Ipv4Addr::LOCALHOST)));
512 }
513 }
514}
515
516#[cfg(test)]
517mod property_tests {
518 use super::*;
519 use proptest::prelude::*;
520
521 fn any_ipv4() -> impl Strategy<Value = Ipv4Addr> {
522 any::<u32>().prop_map(Ipv4Addr::from)
523 }
524
525 fn any_ipv6() -> impl Strategy<Value = Ipv6Addr> {
526 any::<u128>().prop_map(Ipv6Addr::from)
527 }
528
529 proptest::proptest! {
530 #[test]
531 fn v4_membership_invariant(ips in proptest::collection::vec(any_ipv4(), 1..50)) {
532 let mut set = IpSet::new();
533 for &ip in &ips {
534 set.insert(IpAddr::V4(ip));
535 }
536 for ip in ips {
537 prop_assert!(set.contains(&IpAddr::V4(ip)));
538 }
539 }
540
541 #[test]
542 fn v6_membership_invariant(ips in proptest::collection::vec(any_ipv6(), 1..50)) {
543 let mut set = IpSet::new();
544 for &ip in &ips {
545 set.insert(IpAddr::V6(ip));
546 }
547 for ip in ips {
548 prop_assert!(set.contains(&IpAddr::V6(ip)));
549 }
550 }
551
552 #[test]
553 fn order_independence_mixed(
554 ips in proptest::collection::vec(
555 prop_oneof![
556 any_ipv4().prop_map(IpAddr::V4),
557 any_ipv6().prop_map(IpAddr::V6),
558 ],
559 0..50
560 )
561 ) {
562 let mut set1 = IpSet::new();
563 let mut set2 = IpSet::new();
564
565 for &ip in &ips { set1.insert(ip); }
566 let mut ips_rev = ips.clone();
567 ips_rev.reverse();
568 for &ip in &ips_rev { set2.insert(ip); }
569
570 set1.canonicalize();
571 set2.canonicalize();
572 prop_assert_eq!(set1, set2);
573 }
574 }
575}