1mod addrpolicy;
20mod portpolicy;
21mod summary;
22
23use std::fmt;
24use std::ops::RangeInclusive;
25use std::str::FromStr;
26use std::{collections::BTreeSet, fmt::Display};
27use thiserror::Error;
28use tor_basic_utils::iter_join;
29
30pub use addrpolicy::{AddrPolicy, AddrPortPattern, IpPattern};
31pub use portpolicy::PortPolicy;
32pub use summary::{PortPolicies, PortSummaryThresholds};
33
34use crate::NormalItemArgument;
35use crate::parse2::{ArgumentError, ArgumentStream, ItemArgumentParseable};
36
37#[derive(Debug, Error, Clone, PartialEq, Eq)]
39#[non_exhaustive]
40pub enum PolicyError {
41 #[error("Invalid port")]
43 InvalidPort,
44 #[error("Invalid port range")]
46 InvalidRange,
47 #[error("Invalid address")]
49 InvalidAddress,
50 #[error("mask or prefix length with star")]
53 MaskWithStar,
54 #[error("invalid prefix length or mask")]
57 InvalidMask,
58 #[error("Invalid policy")]
60 InvalidPolicy,
61}
62
63#[derive(derive_more::Debug, Clone, Copy, PartialEq, Eq, Hash)]
78#[allow(clippy::exhaustive_structs)]
79#[debug("PortRange({})", &self)]
80pub struct PortRange {
81 lo: u16,
83 hi: u16,
85}
86
87impl PortRange {
88 const fn new_unchecked(lo: u16, hi: u16) -> Self {
91 assert!(lo != 0);
92 assert!(lo <= hi);
93 PortRange { lo, hi }
94 }
95 pub const fn new_all() -> Self {
97 PortRange::new_unchecked(1, 65535)
98 }
99 pub fn new(lo: u16, hi: u16) -> Option<Self> {
105 if lo != 0 && lo <= hi {
106 Some(PortRange { lo, hi })
107 } else {
108 None
109 }
110 }
111 pub fn from_range(r: RangeInclusive<u16>) -> Option<Self> {
115 Self::new(*r.start(), *r.end())
116 }
117 pub fn to_range(self) -> RangeInclusive<u16> {
121 self.lo..=self.hi
122 }
123 pub fn contains(&self, port: u16) -> bool {
125 self.lo <= port && port <= self.hi
126 }
127 pub fn is_all(&self) -> bool {
129 self.lo == 1 && self.hi == 65535
130 }
131
132 fn compare_to_port(&self, port: u16) -> std::cmp::Ordering {
138 use std::cmp::Ordering::*;
139 if port < self.lo {
140 Greater
141 } else if port <= self.hi {
142 Equal
143 } else {
144 Less
145 }
146 }
147}
148
149impl Display for PortRange {
153 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
154 if self.lo == self.hi {
155 write!(f, "{}", self.lo)
156 } else {
157 write!(f, "{}-{}", self.lo, self.hi)
158 }
159 }
160}
161
162impl FromStr for PortRange {
163 type Err = PolicyError;
164 fn from_str(s: &str) -> Result<Self, PolicyError> {
165 let (lo, hi) = match s.split_once('-') {
166 Some((lo, hi)) => (
167 lo.parse::<u16>().map_err(|_| PolicyError::InvalidPort)?,
168 hi.parse::<u16>().map_err(|_| PolicyError::InvalidPort)?,
169 ),
170 None => {
171 let v = s.parse::<u16>().map_err(|_| PolicyError::InvalidPort)?;
173 (v, v)
174 }
175 };
176 PortRange::new(lo, hi).ok_or(PolicyError::InvalidRange)
177 }
178}
179
180impl NormalItemArgument for PortRange {}
181
182#[derive(Debug, Clone, PartialEq, Eq, Hash, Default)]
189struct PortRanges(Vec<PortRange>);
193
194impl PortRanges {
195 fn new() -> Self {
197 Self(Vec::new())
198 }
199
200 fn is_empty(&self) -> bool {
202 self.0.is_empty()
203 }
204
205 fn push_ordered(&mut self, item: PortRange) -> Result<(), PolicyError> {
211 if let Some(prev) = self.0.last() {
212 if prev.hi >= item.lo {
215 return Err(PolicyError::InvalidPolicy);
216 } else if prev.hi == item.lo - 1 {
217 let r = PortRange::new_unchecked(prev.lo, item.hi);
219 self.0.pop();
220 self.0.push(r);
221 return Ok(());
222 }
223 }
224
225 self.0.push(item);
226 Ok(())
227 }
228
229 fn contains(&self, port: u16) -> bool {
235 debug_assert!(self.0.is_sorted_by(|a, b| a.lo < b.lo));
236 self.0
237 .binary_search_by(|range| range.compare_to_port(port))
238 .is_ok()
239 }
240
241 fn inverted(&self) -> PortRanges {
245 let mut prev_hi = 0;
246 let mut new_allowed = Vec::new();
247 for entry in &self.0 {
248 if entry.lo > prev_hi + 1 {
251 new_allowed.push(PortRange::new_unchecked(prev_hi + 1, entry.lo - 1));
252 }
253 prev_hi = entry.hi;
254 }
255 if prev_hi < 65535 {
256 new_allowed.push(PortRange::new_unchecked(prev_hi + 1, 65535));
257 }
258 PortRanges(new_allowed)
259 }
260
261 fn invert(&mut self) {
265 *self = self.inverted();
266 }
267
268 fn iter(&self) -> impl Iterator<Item = &PortRange> + Clone {
270 self.0.iter()
271 }
272
273 fn display(&self) -> Option<impl Display + '_> {
281 struct DisplayWrapper<'r>(&'r PortRanges);
282
283 impl Display for DisplayWrapper<'_> {
284 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
285 write!(f, "{}", iter_join(",", self.0.iter()))
286 }
287 }
288
289 (!self.is_empty()).then_some(DisplayWrapper(self))
290 }
291}
292
293impl FromIterator<u16> for PortRanges {
294 fn from_iter<I: IntoIterator<Item = u16>>(iter: I) -> Self {
295 let ports = iter.into_iter().collect::<BTreeSet<_>>();
297 let mut ports = ports.into_iter().peekable();
298
299 let mut out = Self::new();
300 let mut current_min = None;
301 while let Some(port) = ports.next() {
302 if current_min.is_none() {
303 current_min = Some(port);
304 }
305 if let Some(next_port) = ports.peek().copied() {
306 if next_port != port + 1 {
311 let _ = out.push_ordered(PortRange::new_unchecked(
312 current_min.expect("Don't have min port number"),
313 port,
314 ));
315 current_min = None;
316 }
317 } else {
318 let _ = out.push_ordered(PortRange::new_unchecked(
319 current_min.expect("Don't have min port number"),
320 port,
321 ));
322 }
323 }
324
325 out
326 }
327}
328
329impl FromStr for PortRanges {
330 type Err = PolicyError;
331
332 fn from_str(s: &str) -> Result<Self, Self::Err> {
333 let mut ranges = Self::new();
336 for range in s.split(',') {
337 ranges.push_ordered(range.parse()?)?;
338 }
339 Ok(ranges)
340 }
341}
342
343impl ItemArgumentParseable for PortRanges {
344 fn from_args<'s>(args: &mut ArgumentStream<'s>) -> Result<Self, ArgumentError> {
347 args.next()
348 .map(Self::from_str)
349 .unwrap_or(Ok(Self::new()))
350 .map_err(|_| ArgumentError::Invalid)
351 }
352}
353
354#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash, derive_more::Display, derive_more::FromStr)]
357#[display(rename_all = "lowercase")]
358#[from_str(rename_all = "lowercase")]
359#[allow(clippy::exhaustive_enums)]
360pub enum RuleKind {
361 Accept,
363 Reject,
365}
366
367impl NormalItemArgument for RuleKind {}
368
369#[cfg(test)]
370mod test {
371 #![allow(clippy::bool_assert_comparison)]
373 #![allow(clippy::clone_on_copy)]
374 #![allow(clippy::dbg_macro)]
375 #![allow(clippy::mixed_attributes_style)]
376 #![allow(clippy::print_stderr)]
377 #![allow(clippy::print_stdout)]
378 #![allow(clippy::single_char_pattern)]
379 #![allow(clippy::unwrap_used)]
380 #![allow(clippy::unchecked_time_subtraction)]
381 #![allow(clippy::useless_vec)]
382 #![allow(clippy::needless_pass_by_value)]
383 #![allow(clippy::string_slice)] use super::*;
386 use crate::Result;
387 use crate::parse2::{self, ParseInput};
388
389 #[test]
390 fn parse_portrange() -> Result<()> {
391 assert_eq!(
392 "1-100".parse::<PortRange>()?,
393 PortRange::new(1, 100).unwrap()
394 );
395 assert_eq!(
396 "01-100".parse::<PortRange>()?,
397 PortRange::new(1, 100).unwrap()
398 );
399 assert_eq!("1-65535".parse::<PortRange>()?, PortRange::new_all());
400 assert_eq!(
401 "10-30".parse::<PortRange>()?,
402 PortRange::new(10, 30).unwrap()
403 );
404 assert_eq!(
405 "9001".parse::<PortRange>()?,
406 PortRange::new(9001, 9001).unwrap()
407 );
408 assert_eq!(
409 "9001-9001".parse::<PortRange>()?,
410 PortRange::new(9001, 9001).unwrap()
411 );
412
413 assert!("hello".parse::<PortRange>().is_err());
414 assert!("0".parse::<PortRange>().is_err());
415 assert!("65536".parse::<PortRange>().is_err());
416 assert!("65537".parse::<PortRange>().is_err());
417 assert!("1-2-3".parse::<PortRange>().is_err());
418 assert!("10-5".parse::<PortRange>().is_err());
419 assert!("1-".parse::<PortRange>().is_err());
420 assert!("-2".parse::<PortRange>().is_err());
421 assert!("-".parse::<PortRange>().is_err());
422 assert!("*".parse::<PortRange>().is_err());
423 Ok(())
424 }
425
426 #[test]
427 fn pr_manip() {
428 assert!(PortRange::new_all().is_all());
429 assert!(!PortRange::new(2, 65535).unwrap().is_all());
430
431 assert!(PortRange::new_all().contains(1));
432 assert!(PortRange::new_all().contains(65535));
433 assert!(PortRange::new_all().contains(7777));
434
435 assert!(PortRange::new(20, 30).unwrap().contains(20));
436 assert!(PortRange::new(20, 30).unwrap().contains(25));
437 assert!(PortRange::new(20, 30).unwrap().contains(30));
438 assert!(!PortRange::new(20, 30).unwrap().contains(19));
439 assert!(!PortRange::new(20, 30).unwrap().contains(31));
440
441 use std::cmp::Ordering::*;
442 assert_eq!(PortRange::new(20, 30).unwrap().compare_to_port(7), Greater);
443 assert_eq!(PortRange::new(20, 30).unwrap().compare_to_port(20), Equal);
444 assert_eq!(PortRange::new(20, 30).unwrap().compare_to_port(25), Equal);
445 assert_eq!(PortRange::new(20, 30).unwrap().compare_to_port(30), Equal);
446 assert_eq!(PortRange::new(20, 30).unwrap().compare_to_port(100), Less);
447 }
448
449 #[test]
450 fn pr_fmt() {
451 fn chk(a: u16, b: u16, s: &str) {
452 let pr = PortRange::new(a, b).unwrap();
453 assert_eq!(format!("{}", pr), s);
454 }
455
456 chk(1, 65535, "1-65535");
457 chk(10, 20, "10-20");
458 chk(20, 20, "20");
459 }
460
461 #[test]
462 fn port_ranges() {
463 const INPUT: &str = "22,80,443,8000-9000,9002";
464 let ranges = PortRanges::from_str(INPUT).unwrap();
465 assert_eq!(
466 ranges.0,
467 [
468 PortRange::new(22, 22).unwrap(),
469 PortRange::new(80, 80).unwrap(),
470 PortRange::new(443, 443).unwrap(),
471 PortRange::new(8000, 9000).unwrap(),
472 PortRange::new(9002, 9002).unwrap(),
473 ]
474 );
475 assert!(ranges.contains(22));
476 assert!(ranges.contains(80));
477 assert!(ranges.contains(443));
478 assert!(ranges.contains(8000));
479 assert!(ranges.contains(8500));
480 assert!(ranges.contains(9000));
481 assert!(!ranges.contains(9001));
482 assert!(ranges.contains(9002));
483
484 let mut ranges_inverse = ranges.clone();
485 ranges_inverse.invert();
486 assert_eq!(
487 ranges_inverse.0,
488 [
489 PortRange::new(1, 21).unwrap(),
490 PortRange::new(23, 79).unwrap(),
491 PortRange::new(81, 442).unwrap(),
492 PortRange::new(444, 7999).unwrap(),
493 PortRange::new(9001, 9001).unwrap(),
494 PortRange::new(9003, 65535).unwrap(),
495 ]
496 );
497
498 #[derive(derive_deftly::Deftly)]
499 #[derive_deftly(NetdocParseable)]
500 struct Dummy {
501 #[deftly(netdoc(single_arg))]
502 dummy: PortRanges,
503 }
504 let ranges2 =
505 parse2::parse_netdoc::<Dummy>(&ParseInput::new(&format!("dummy {INPUT}\n"), ""))
506 .unwrap();
507 assert_eq!(ranges, ranges2.dummy);
508 }
509}