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!(!self.v4_dirty && !self.v6_dirty, "IpSet must be canonicalized before calling contains_canonical");
218 match ip {
219 IpAddr::V4(v4) => {
220 let target = u32::from(*v4);
221 self.v4.binary_search_by(|range| {
222 let start = u32::from(range.start_addr);
223 let end = u32::from(range.end_addr);
224 if target < start { std::cmp::Ordering::Greater }
225 else if target > end { std::cmp::Ordering::Less }
226 else { std::cmp::Ordering::Equal }
227 }).is_ok()
228 }
229 IpAddr::V6(v6) => {
230 let target = u128::from(*v6);
231 self.v6.binary_search_by(|range| {
232 let start = u128::from(range.start_addr);
233 let end = u128::from(range.end_addr);
234 if target < start { std::cmp::Ordering::Greater }
235 else if target > end { std::cmp::Ordering::Less }
236 else { std::cmp::Ordering::Equal }
237 }).is_ok()
238 }
239 }
240 }
241
242 pub fn len_canonical(&self) -> u128 {
248 debug_assert!(!self.v4_dirty && !self.v6_dirty, "IpSet must be canonicalized before calling len_canonical");
249 let v4_len: u128 = self.v4.iter().map(|r| r.len() as u128).sum();
250 let v6_len: u128 = self.v6.iter().map(|r| r.len()).sum();
251 v4_len + v6_len
252 }
253
254 pub fn v4(&self) -> &[Ipv4Range] {
256 &self.v4
257 }
258
259 pub fn v6(&self) -> &[Ipv6Range] {
261 &self.v6
262 }
263}
264
265impl IntoIterator for IpSet {
270 type Item = IpAddr;
271 type IntoIter = Box<dyn Iterator<Item = IpAddr> + Send>;
272
273 fn into_iter(mut self) -> Self::IntoIter {
275 self.canonicalize();
276 let v4_iter = self.v4.into_iter().flat_map(|range| {
277 let start: u32 = range.start_addr.into();
278 let end: u32 = range.end_addr.into();
279 (start..=end).map(|ip| IpAddr::V4(Ipv4Addr::from(ip)))
280 });
281
282 let v6_iter = self.v6.into_iter().flat_map(|range| {
283 let start: u128 = range.start_addr.into();
284 let end: u128 = range.end_addr.into();
285 (start..=end).map(|ip| IpAddr::V6(Ipv6Addr::from(ip)))
286 });
287
288 Box::new(v4_iter.chain(v6_iter))
289 }
290}
291
292impl Extend<IpAddr> for IpSet {
293 fn extend<T: IntoIterator<Item = IpAddr>>(&mut self, iter: T) {
294 for ip in iter {
295 match ip {
296 IpAddr::V4(v4) => self.v4.push(Ipv4Range::new(v4, v4).unwrap()),
297 IpAddr::V6(v6) => self.v6.push(Ipv6Range::new(v6, v6).unwrap()),
298 }
299 }
300 self.v4_dirty = true;
301 self.v6_dirty = true;
302 }
303}
304
305impl FromIterator<IpAddr> for IpSet {
306 fn from_iter<I: IntoIterator<Item = IpAddr>>(iter: I) -> Self {
307 let mut set = IpSet::new();
308 set.extend(iter);
309 set.canonicalize();
310 set
311 }
312}
313
314impl FromIterator<IpRange> for IpSet {
315 fn from_iter<I: IntoIterator<Item = IpRange>>(iter: I) -> Self {
316 let mut set = IpSet::new();
317 for range in iter {
318 match range {
319 IpRange::V4(r) => set.v4.push(r),
320 IpRange::V6(r) => set.v6.push(r),
321 }
322 }
323 set.v4_dirty = true;
324 set.v6_dirty = true;
325 set.canonicalize();
326 set
327 }
328}
329
330impl FromIterator<IpSet> for IpSet {
331 fn from_iter<I: IntoIterator<Item = IpSet>>(iter: I) -> Self {
332 let mut master = IpSet::new();
333 for set in iter {
334 master.v4.extend(set.v4);
335 master.v6.extend(set.v6);
336 }
337 master.v4_dirty = true;
338 master.v6_dirty = true;
339 master.canonicalize();
340 master
341 }
342}
343
344impl From<IpAddr> for IpSet {
345 fn from(ip: IpAddr) -> Self {
346 let mut set = Self::new();
347 set.insert(ip);
348 set
349 }
350}
351
352impl From<IpRange> for IpSet {
353 fn from(range: IpRange) -> Self {
354 let mut set = Self::new();
355 set.insert_range(range);
356 set
357 }
358}
359
360impl TryFrom<&str> for IpSet {
361 type Error = IpSetError;
362 fn try_from(value: &str) -> Result<Self, Self::Error> {
363 let mut set = IpSet::new();
364 for part in value.split([',', ' ']).filter(|part| !part.trim().is_empty()) {
365 let range = part.parse::<IpRange>()?;
366 set.insert_range(range);
367 }
368 set.canonicalize();
369 Ok(set)
370 }
371}
372
373impl FromStr for IpSet {
374 type Err = IpSetError;
375 fn from_str(s: &str) -> Result<Self, Self::Err> {
376 Self::try_from(s)
377 }
378}
379
380#[cfg(test)]
390mod tests {
391 use super::*;
392
393 #[test]
394 fn lazy_merging_v4() {
395 let mut set = IpSet::new();
396 set.insert(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)));
397 set.insert(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 2)));
398
399 assert_eq!(set.v4.len(), 2);
401 assert!(set.v4_dirty);
402
403 set.canonicalize();
405 assert_eq!(set.len(), 2);
406 assert!(!set.v4_dirty);
407 assert_eq!(set.v4.len(), 1);
408 }
409
410 #[test]
411 fn set_battle_test_overlaps() {
412 let mut set = IpSet::new();
413 set.insert_range("10.0.0.10-10.0.0.20".parse().unwrap());
415 set.insert_range("10.0.0.5-10.0.0.15".parse().unwrap());
417 set.insert_range("10.0.0.15-10.0.0.25".parse().unwrap());
419 set.insert_range("10.0.0.30-10.0.0.40".parse().unwrap());
421 set.insert_range("10.0.0.0-10.0.0.50".parse().unwrap());
423
424 set.canonicalize();
425 assert_eq!(set.len(), 51);
426 assert_eq!(set.v4().len(), 1);
427 }
428
429 #[test]
430 fn set_ipv6_adjacency_boundary() {
431 let mut set = IpSet::new();
432 let max_v6 = Ipv6Addr::from(u128::MAX);
434 let max_minus_1 = Ipv6Addr::from(u128::MAX - 1);
435
436 set.insert(IpAddr::V6(max_minus_1));
437 set.insert(IpAddr::V6(max_v6));
438
439 set.canonicalize();
440 assert_eq!(set.len(), 2);
441 assert_eq!(set.v6().len(), 1);
442 }
443
444 #[test]
445 fn iteration_is_lazy_safe() {
446 let mut set = IpSet::new();
447 set.insert(IpAddr::V4(Ipv4Addr::from(1)));
448 set.insert(IpAddr::V4(Ipv4Addr::from(2)));
449
450 set.canonicalize();
451 let ips: Vec<IpAddr> = set.iter().collect();
452 assert_eq!(ips.len(), 2);
453 assert!(!set.v4_dirty);
454 }
455
456 #[test]
457 fn empty_set_canonical_is_fine() {
458 let mut set = IpSet::new();
459 set.canonicalize();
460 assert_eq!(set.len_canonical(), 0);
461 assert!(set.v4().is_empty());
462 }
463
464 #[test]
465 fn from_str_mixed_advanced() {
466 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();
467 assert_eq!(set.len(), 4);
469 }
470
471 #[test]
472 fn bulk_extend_efficiency() {
473 let mut set = IpSet::new();
474 let ips = (0..100).map(|i| IpAddr::V4(Ipv4Addr::from(i)));
475 set.extend(ips);
476
477 assert_eq!(set.v4.len(), 100);
478 set.canonicalize();
479 assert_eq!(set.v4.len(), 1);
480 assert_eq!(set.len(), 100);
481 }
482
483 #[test]
484 fn canonical_queries_panics_in_debug() {
485 #[cfg(debug_assertions)]
486 {
487 let set = IpSet::from_iter(vec![IpAddr::V4(Ipv4Addr::LOCALHOST)]);
488 assert!(!set.v4_dirty);
490 assert!(set.contains_canonical(&IpAddr::V4(Ipv4Addr::LOCALHOST)));
491 }
492 }
493}
494
495#[cfg(test)]
496mod property_tests {
497 use super::*;
498 use proptest::prelude::*;
499
500 fn any_ipv4() -> impl Strategy<Value = Ipv4Addr> {
501 any::<u32>().prop_map(Ipv4Addr::from)
502 }
503
504 fn any_ipv6() -> impl Strategy<Value = Ipv6Addr> {
505 any::<u128>().prop_map(Ipv6Addr::from)
506 }
507
508 proptest::proptest! {
509 #[test]
510 fn v4_membership_invariant(ips in proptest::collection::vec(any_ipv4(), 1..50)) {
511 let mut set = IpSet::new();
512 for &ip in &ips {
513 set.insert(IpAddr::V4(ip));
514 }
515 for ip in ips {
516 prop_assert!(set.contains(&IpAddr::V4(ip)));
517 }
518 }
519
520 #[test]
521 fn v6_membership_invariant(ips in proptest::collection::vec(any_ipv6(), 1..50)) {
522 let mut set = IpSet::new();
523 for &ip in &ips {
524 set.insert(IpAddr::V6(ip));
525 }
526 for ip in ips {
527 prop_assert!(set.contains(&IpAddr::V6(ip)));
528 }
529 }
530
531 #[test]
532 fn order_independence_mixed(
533 ips in proptest::collection::vec(
534 prop_oneof![
535 any_ipv4().prop_map(IpAddr::V4),
536 any_ipv6().prop_map(IpAddr::V6),
537 ],
538 0..50
539 )
540 ) {
541 let mut set1 = IpSet::new();
542 let mut set2 = IpSet::new();
543
544 for &ip in &ips { set1.insert(ip); }
545 let mut ips_rev = ips.clone();
546 ips_rev.reverse();
547 for &ip in &ips_rev { set2.insert(ip); }
548
549 set1.canonicalize();
550 set2.canonicalize();
551 prop_assert_eq!(set1, set2);
552 }
553 }
554}