1use std::fmt::Display;
5use std::net::{IpAddr, SocketAddr};
6use std::str::FromStr;
7
8use itertools::chain;
9
10use crate::NormalItemArgument;
11use crate::encode::NetdocEncodableFields;
12use crate::parse2::{
13 ErrorProblem as EP, ItemArgumentParseable, ItemStream, KeywordRef, NetdocParseableFields,
14 UnparsedItem,
15};
16
17use ipnet::IpNet;
18
19use super::{PolicyError, PortRange, RuleKind};
20
21#[derive(Clone, Debug, Default, PartialEq, Eq)]
55pub struct AddrPolicy {
56 rules: Vec<AddrPolicyRule>,
62}
63
64impl AddrPolicy {
65 pub fn allows(&self, addr: &IpAddr, port: u16) -> Option<RuleKind> {
72 self.rules
73 .iter()
74 .find(|rule| rule.pattern.matches(addr, port))
75 .map(|AddrPolicyRule { kind, .. }| *kind)
76 }
77
78 pub fn allows_sockaddr(&self, addr: &SocketAddr) -> Option<RuleKind> {
80 self.allows(&addr.ip(), addr.port())
81 }
82
83 pub fn new() -> Self {
85 AddrPolicy::default()
86 }
87
88 pub fn push(&mut self, kind: RuleKind, pattern: AddrPortPattern) {
96 self.rules.push(AddrPolicyRule { kind, pattern });
97 }
98
99 pub fn rules(&self) -> impl DoubleEndedIterator<Item = (RuleKind, AddrPortPattern)> + '_ {
101 self.rules
102 .iter()
103 .map(|rule| (rule.kind, rule.pattern.clone()))
104 }
105}
106
107impl NetdocParseableFields for AddrPolicy {
108 type Accumulator = AddrPolicy;
109
110 fn is_item_keyword(kw: KeywordRef<'_>) -> bool {
111 matches!(kw.as_str(), "accept" | "reject")
112 }
113
114 fn accumulate_item(acc: &mut Self::Accumulator, mut item: UnparsedItem<'_>) -> Result<(), EP> {
115 let rule = RuleKind::from_str(item.keyword().as_str())
118 .map_err(|_| EP::Internal("accept/reject not a RuleKind?"))?;
119 let args = item.args_mut();
120 let pattern =
121 AddrPortPattern::from_args(args).map_err(args.error_handler("accept/reject"))?;
122 acc.push(rule, pattern);
123 Ok(())
124 }
125
126 fn finish(acc: Self::Accumulator, _: &ItemStream) -> Result<Self, EP> {
127 Ok(acc)
128 }
129}
130
131impl NetdocEncodableFields for AddrPolicy {
132 fn encode_fields(&self, out: &mut crate::encode::NetdocEncoder) -> Result<(), tor_error::Bug> {
133 const ALL: AddrPortPattern = AddrPortPattern::new_all();
141 const DEFAULT_DENY: AddrPolicyRule = AddrPolicyRule {
142 kind: RuleKind::Reject,
143 pattern: ALL,
144 };
145
146 let default_deny = match self.rules.last() {
148 Some(AddrPolicyRule {
150 kind: _,
151 pattern: ALL,
152 }) => None,
153 _ => Some(&DEFAULT_DENY),
155 };
156
157 for rule in chain!(&self.rules, default_deny) {
158 out.push_raw_string(&format_args!("{} {}\n", rule.kind, rule.pattern));
159 }
160 Ok(())
161 }
162}
163
164#[derive(Clone, Debug, PartialEq, Eq)]
168struct AddrPolicyRule {
169 kind: RuleKind,
171 pattern: AddrPortPattern,
173}
174
175#[derive(Clone, Debug, Eq, PartialEq, Hash)] #[derive(serde_with::SerializeDisplay, serde_with::DeserializeFromStr)]
211#[allow(clippy::exhaustive_structs)]
212pub struct AddrPortPattern {
213 pub addrs: IpPattern,
215 pub ports: PortRange,
217}
218
219impl AddrPortPattern {
220 pub fn new(addrs: IpPattern, ports: PortRange) -> Self {
222 Self { addrs, ports }
223 }
224
225 pub const fn new_all() -> Self {
227 Self {
228 addrs: IpPattern::All,
229 ports: PortRange::new_all(),
230 }
231 }
232
233 pub fn matches(&self, addr: &IpAddr, port: u16) -> bool {
235 self.addrs.matches(addr) && self.ports.contains(port)
236 }
237 pub fn matches_sockaddr(&self, addr: &SocketAddr) -> bool {
239 self.matches(&addr.ip(), addr.port())
240 }
241}
242
243impl Display for AddrPortPattern {
244 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
245 if self.ports.is_all() {
246 write!(f, "{}:*", self.addrs)
247 } else {
248 write!(f, "{}:{}", self.addrs, self.ports)
249 }
250 }
251}
252
253impl FromStr for AddrPortPattern {
254 type Err = PolicyError;
255 fn from_str(s: &str) -> Result<Self, PolicyError> {
256 let (addrs, ports_s) = s.rsplit_once(':').ok_or(PolicyError::InvalidPolicy)?;
257 let addrs: IpPattern = addrs.parse()?;
258 let ports: PortRange = if ports_s == "*" {
259 PortRange::new_all()
260 } else {
261 ports_s.parse()?
262 };
263
264 Ok(AddrPortPattern { addrs, ports })
265 }
266}
267
268impl NormalItemArgument for AddrPortPattern {}
269
270#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash, derive_more::From)]
272#[allow(clippy::exhaustive_enums)]
275pub enum IpPattern {
276 All,
283 Net(#[from] IpNet),
288}
289
290impl IpPattern {
291 pub fn from_addr_and_prefix_len(addr: IpAddr, prefix_len: u8) -> Result<Self, PolicyError> {
293 IpNet::new(addr, prefix_len)
294 .map(IpPattern::Net)
295 .map_err(|_: ipnet::PrefixLenError| PolicyError::InvalidMask)
296 }
297
298 pub fn matches(&self, addr: &IpAddr) -> bool {
300 match self {
301 IpPattern::All => true,
302 IpPattern::Net(n) => n.contains(addr),
303 }
304 }
305}
306
307impl Display for IpPattern {
308 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
309 use IpPattern::*;
310 match self {
311 All => write!(f, "*"),
312 Net(IpNet::V4(n)) if n.prefix_len() == 32 => write!(f, "{}", n.addr()),
314 Net(IpNet::V4(n)) => write!(f, "{}", n),
315 Net(IpNet::V6(n)) if n.prefix_len() == 128 => write!(f, "[{}]", n.addr()),
317 Net(IpNet::V6(n)) => write!(f, "[{}]/{}", n.addr(), n.prefix_len()),
318 }
319 }
320}
321
322fn parse_addr(mut s: &str) -> Result<IpAddr, PolicyError> {
325 let trimmed = s.strip_prefix('[').and_then(|s| s.strip_suffix(']'));
326 if let Some(trimmed) = trimmed {
327 s = trimmed;
328 }
329 let addr: IpAddr = s.parse().map_err(|_| PolicyError::InvalidAddress)?;
330 if addr.is_ipv6() != trimmed.is_some() {
331 return Err(PolicyError::InvalidAddress);
332 }
333 Ok(addr)
334}
335
336impl FromStr for IpPattern {
337 type Err = PolicyError;
338 fn from_str(s: &str) -> Result<Self, PolicyError> {
339 let (ip_s, plen_s) = match s.split_once('/') {
340 Some((ip_s, plen_s)) => (ip_s, Some(plen_s)),
341 None => (s, None),
342 };
343 match (ip_s, plen_s) {
344 ("*", Some(_)) => Err(PolicyError::MaskWithStar),
345 ("*", None) => Ok(IpPattern::All),
346 (s, Some(m)) => {
347 let a: IpAddr = parse_addr(s)?;
348 let m: u8 = m.parse().map_err(|_| PolicyError::InvalidMask)?;
349 IpPattern::from_addr_and_prefix_len(a, m)
350 }
351 (s, None) => {
352 let a: IpAddr = parse_addr(s)?;
353 let m = if a.is_ipv4() { 32 } else { 128 };
354 IpPattern::from_addr_and_prefix_len(a, m)
355 }
356 }
357 }
358}
359
360#[cfg(test)]
361mod test {
362 #![allow(clippy::bool_assert_comparison)]
364 #![allow(clippy::clone_on_copy)]
365 #![allow(clippy::dbg_macro)]
366 #![allow(clippy::mixed_attributes_style)]
367 #![allow(clippy::print_stderr)]
368 #![allow(clippy::print_stdout)]
369 #![allow(clippy::single_char_pattern)]
370 #![allow(clippy::unwrap_used)]
371 #![allow(clippy::unchecked_time_subtraction)]
372 #![allow(clippy::useless_vec)]
373 #![allow(clippy::needless_pass_by_value)]
374 #![allow(clippy::string_slice)] use crate::encode::{NetdocEncodable, NetdocEncoder};
377
378 use super::*;
379
380 #[test]
381 fn test_roundtrip_rules() {
382 fn check2(inp: &str, outp: &str) {
383 let policy = inp.parse::<AddrPortPattern>().expect(inp);
384 assert_eq!(format!("{}", policy), outp);
385 }
386 let check = |inp| check2(inp, inp);
387
388 check2("127.0.0.2/32:77-10000", "127.0.0.2:77-10000");
389 check2("127.0.0.2/32:*", "127.0.0.2:*");
390 check("127.0.0.0/16:9-100");
391 check("127.0.0.0/0:443");
392 check("*:443");
393 check("[::1]:443");
394 check("[ffaa::]/16:80");
395 check2("[ffaa::77]/128:80", "[ffaa::77]:80");
396 check("[::]/0:443");
397
398 check("127.0.0.1/8:443");
401 check("[::1]/8:443");
402 }
403
404 #[test]
405 fn test_bad_rules() {
406 fn check(s: &str) {
407 let _: PolicyError = s.parse::<AddrPortPattern>().expect_err(s);
408 }
409
410 check("marzipan:80");
411 check("1.2.3.4:90-80");
412 check("1.2.3.4/100:8888");
413 check("[1.2.3.4]/16:80");
414 check("[::1]/130:8888");
415 }
416
417 #[test]
418 fn test_rule_matches() {
419 fn check(addr: &str, yes: &[&str], no: &[&str]) {
420 use std::net::SocketAddr;
421 let policy = addr.parse::<AddrPortPattern>().unwrap();
422 for s in yes {
423 let sa = s.parse::<SocketAddr>().unwrap();
424 assert!(policy.matches_sockaddr(&sa));
425 }
426 for s in no {
427 let sa = s.parse::<SocketAddr>().unwrap();
428 assert!(!policy.matches_sockaddr(&sa));
429 }
430 }
431
432 check(
433 "1.2.3.4/16:80",
434 &["1.2.3.4:80", "1.2.44.55:80"],
435 &["9.9.9.9:80", "1.3.3.4:80", "1.2.3.4:81"],
436 );
437 check(
438 "*:443-8000",
439 &["1.2.3.4:443", "[::1]:500"],
440 &["9.0.0.0:80", "[::1]:80"],
441 );
442 check(
443 "[face::]/8:80",
444 &["[fab0::7]:80"],
445 &["[dd00::]:80", "[face::7]:443"],
446 );
447
448 check("0.0.0.0/0:*", &["127.0.0.1:80"], &["[f00b::]:80"]);
449 check("[::]/0:*", &["[f00b::]:80"], &["127.0.0.1:80"]);
450 }
451
452 #[test]
453 fn test_policy_matches() -> Result<(), PolicyError> {
454 let mut policy = AddrPolicy::default();
455 policy.push(RuleKind::Accept, "*:443".parse()?);
456 policy.push(RuleKind::Accept, "[::1]:80".parse()?);
457 policy.push(RuleKind::Reject, "*:80".parse()?);
458
459 let policy = policy; assert_eq!(
461 policy.allows_sockaddr(&"[::6]:443".parse().unwrap()),
462 Some(RuleKind::Accept)
463 );
464 assert_eq!(
465 policy.allows_sockaddr(&"127.0.0.1:443".parse().unwrap()),
466 Some(RuleKind::Accept)
467 );
468 assert_eq!(
469 policy.allows_sockaddr(&"[::1]:80".parse().unwrap()),
470 Some(RuleKind::Accept)
471 );
472 assert_eq!(
473 policy.allows_sockaddr(&"[::2]:80".parse().unwrap()),
474 Some(RuleKind::Reject)
475 );
476 assert_eq!(
477 policy.allows_sockaddr(&"127.0.0.1:80".parse().unwrap()),
478 Some(RuleKind::Reject)
479 );
480 assert_eq!(
481 policy.allows_sockaddr(&"127.0.0.1:66".parse().unwrap()),
482 None
483 );
484 Ok(())
485 }
486
487 #[test]
488 fn serde() {
489 #[derive(Clone, Debug, serde::Serialize, serde::Deserialize, Eq, PartialEq)]
490 struct X {
491 p1: AddrPortPattern,
492 p2: AddrPortPattern,
493 }
494
495 let x = X {
496 p1: "127.0.0.1/8:9-10".parse().unwrap(),
497 p2: "*:80".parse().unwrap(),
498 };
499
500 let encoded = serde_json::to_string(&x).unwrap();
501 let expected = r#"{"p1":"127.0.0.1/8:9-10","p2":"*:80"}"#;
502 let x2: X = serde_json::from_str(&encoded).unwrap();
503 let x3: X = serde_json::from_str(expected).unwrap();
504 assert_eq!(&x2, &x3);
505 assert_eq!(&x2, &x);
506 }
507
508 #[test]
509 fn parse2() {
510 use crate::parse2::{self, ParseInput};
511 use derive_deftly::Deftly;
512
513 const RULES: &str = "\
514 intro\n\
515 reject *:25\n\
516 reject 127.0.0.0/8:*\n\
517 reject 192.168.0.0/16:*\n\
518 accept *:80\n\
519 accept *:443\n\
520 accept *:9000-65535\n\
521 reject *:*\n";
522
523 #[derive(Deftly)]
524 #[derive_deftly(NetdocParseable, NetdocEncodable)]
525 struct Wrapper {
526 #[allow(dead_code)]
527 intro: (),
528 #[deftly(netdoc(flatten))]
529 ipv4_policy: AddrPolicy,
530 }
531
532 let wrapper = parse2::parse_netdoc::<Wrapper>(&ParseInput::new(RULES, "")).unwrap();
533 let ap = wrapper.ipv4_policy.clone();
534
535 assert_eq!(
536 ap.allows_sockaddr(&"1.1.1.1:80".parse().unwrap()),
537 Some(RuleKind::Accept)
538 );
539 assert_eq!(
540 ap.allows_sockaddr(&"1.1.1.1:443".parse().unwrap()),
541 Some(RuleKind::Accept)
542 );
543 assert_eq!(
544 ap.allows_sockaddr(&"1.1.1.1:9005".parse().unwrap()),
545 Some(RuleKind::Accept)
546 );
547
548 assert_eq!(
549 ap.allows_sockaddr(&"1.1.1.1:25".parse().unwrap()),
550 Some(RuleKind::Reject)
551 );
552 assert_eq!(
553 ap.allows_sockaddr(&"127.0.0.1:80".parse().unwrap()),
554 Some(RuleKind::Reject)
555 );
556 assert_eq!(
557 ap.allows_sockaddr(&"1.1.1.1:70".parse().unwrap()),
558 Some(RuleKind::Reject)
559 );
560
561 let mut enc = NetdocEncoder::default();
563 wrapper.encode_unsigned(&mut enc).unwrap();
564 assert_eq!(RULES, enc.finish().unwrap());
565
566 let accept_all = {
568 let mut ap = AddrPolicy::new();
569 ap.push(RuleKind::Accept, AddrPortPattern::new_all());
570 ap
571 };
572 let reject_all = {
573 let mut ap = AddrPolicy::new();
574 ap.push(RuleKind::Reject, AddrPortPattern::new_all());
575 ap
576 };
577
578 let tests = [
579 (accept_all.clone(), "accept *:*"), (reject_all.clone(), "reject *:*"), (AddrPolicy::new(), "reject *:*"), (AddrPolicy::default(), "reject *:*"), ];
584 for (input, expected_tail) in tests {
585 let mut enc = NetdocEncoder::default();
586 Wrapper {
587 intro: (),
588 ipv4_policy: input,
589 }
590 .encode_unsigned(&mut enc)
591 .unwrap();
592 assert_eq!(expected_tail, enc.finish().unwrap().lines().last().unwrap());
593 }
594 }
595}