1use std::net::IpAddr;
34use std::sync::Arc;
35use std::time::Duration;
36
37use async_trait::async_trait;
38use regex::Regex;
39use tracing::{debug, info};
40
41use super::policy::{Check, StageSet, Verdict};
42use super::{
43 ConnectionContext, IdentifierContext, ListVerdict, canonical, check_lists, compile_matchers,
44};
45use crate::config::DnsConfig;
46use crate::dns::{HickoryResolver, Resolver, resolver_addr};
47
48#[derive(Debug, Clone)]
50pub struct Settings {
51 pub require_forward_confirm: bool,
52 pub allow: Vec<String>,
53 pub deny: Vec<String>,
54 pub allow_regex: Vec<String>,
55 pub deny_regex: Vec<String>,
56 pub timeout_ms: u64,
57}
58
59impl Default for Settings {
60 fn default() -> Self {
61 Self {
62 require_forward_confirm: true,
63 allow: Vec::new(),
64 deny: Vec::new(),
65 allow_regex: Vec::new(),
66 deny_regex: Vec::new(),
67 timeout_ms: 2000,
68 }
69 }
70}
71
72pub struct ClientHasValidReverseDns {
75 resolver: Arc<dyn Resolver>,
76 require_forward_confirm: bool,
77 allow: Vec<Regex>,
78 deny: Vec<Regex>,
79 timeout: Duration,
80}
81
82impl std::fmt::Debug for ClientHasValidReverseDns {
83 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
85 formatter
86 .debug_struct("ClientHasValidReverseDns")
87 .field("require_forward_confirm", &self.require_forward_confirm)
88 .field("allow", &self.allow)
89 .field("deny", &self.deny)
90 .field("timeout", &self.timeout)
91 .finish_non_exhaustive()
92 }
93}
94
95impl ClientHasValidReverseDns {
96 pub fn from_settings(name: &str, settings: &Settings, dns: &DnsConfig) -> anyhow::Result<Self> {
104 let resolver: Arc<dyn Resolver> = Arc::new(match resolver_addr(dns)? {
105 Some(addr) => HickoryResolver::from_address(addr)
106 .map_err(|error| anyhow::anyhow!("filter.check.{name}: {error}"))?,
107 None => HickoryResolver::from_system()
108 .map_err(|error| anyhow::anyhow!("filter.check.{name}: {error}"))?,
109 });
110 let check = Self::with_resolver(name, settings, resolver)?;
111 info!(
112 event = "filter_reverse_dns_loaded",
113 outcome = "success",
114 check = name,
115 require_forward_confirm = settings.require_forward_confirm,
116 timeout_ms = settings.timeout_ms,
117 );
118 Ok(check)
119 }
120
121 pub fn with_resolver(
123 name: &str,
124 settings: &Settings,
125 resolver: Arc<dyn Resolver>,
126 ) -> anyhow::Result<Self> {
127 Ok(Self {
128 resolver,
129 require_forward_confirm: settings.require_forward_confirm,
130 allow: compile_matchers(&settings.allow, &settings.allow_regex, name, "allow")?,
131 deny: compile_matchers(&settings.deny, &settings.deny_regex, name, "deny")?,
132 timeout: Duration::from_millis(settings.timeout_ms),
133 })
134 }
135
136 async fn resolve_hostname(&self, client_ip: IpAddr) -> Result<String, Verdict> {
138 let names = self
139 .resolver
140 .reverse(client_ip)
141 .await
142 .map_err(|error| Verdict::Undecided(format!("PTR lookup failed: {error}")))?;
143
144 if names.is_empty() {
145 return Err(Verdict::Fail(format!("no PTR record for {client_ip}")));
146 }
147
148 if let Some(denied) = names.iter().find(|name| {
160 check_lists(&[], &self.deny, |pattern: &Regex| pattern.is_match(name))
161 == ListVerdict::Denied
162 }) {
163 return Err(Verdict::Fail(format!("hostname {denied} is denied")));
164 }
165
166 let mut last_refusal = None;
170 for name in &names {
171 match self.vet(client_ip, name).await {
172 Ok(()) => return Ok(name.clone()),
173 Err(refusal) => {
174 debug!(event = "filter_reverse_dns_candidate_refused", outcome = "failure", name, reason = ?refusal);
175 last_refusal = Some(refusal);
176 }
177 }
178 }
179
180 Err(last_refusal
181 .unwrap_or_else(|| Verdict::Fail(format!("no acceptable PTR name for {client_ip}"))))
182 }
183
184 async fn vet(&self, client_ip: IpAddr, name: &str) -> Result<(), Verdict> {
187 if check_lists(&self.allow, &[], |pattern: &Regex| pattern.is_match(name))
189 == ListVerdict::NotAllowed
190 {
191 return Err(Verdict::Fail(format!("hostname {name} is not allowed")));
192 }
193
194 if self.require_forward_confirm {
195 let addresses =
196 self.resolver.forward(name).await.map_err(|error| {
197 Verdict::Undecided(format!("forward lookup failed: {error}"))
198 })?;
199
200 if !addresses.iter().any(|addr| canonical(*addr) == client_ip) {
201 return Err(Verdict::Fail(format!(
202 "hostname {name} does not resolve back to {client_ip}"
203 )));
204 }
205 }
206
207 Ok(())
208 }
209
210 async fn decide(&self, client_ip: Option<IpAddr>) -> Verdict {
213 let client_ip = match super::require_client_ip(client_ip) {
216 Ok(client_ip) => client_ip,
217 Err(verdict) => return verdict,
218 };
219
220 match tokio::time::timeout(self.timeout, self.resolve_hostname(client_ip)).await {
223 Ok(Ok(hostname)) => {
224 debug!(event = "filter_reverse_dns_accepted", outcome = "success", client_ip = %client_ip, hostname);
225 Verdict::Pass
226 }
227 Ok(Err(verdict)) => verdict,
228 Err(_) => Verdict::Undecided(format!(
229 "reverse DNS for {client_ip} timed out after {}ms",
230 self.timeout.as_millis()
231 )),
232 }
233 }
234}
235
236#[async_trait]
237impl Check for ClientHasValidReverseDns {
238 fn kind(&self) -> &'static str {
239 "reverse_dns"
240 }
241
242 fn stages(&self) -> StageSet {
246 StageSet::connection_only()
247 }
248
249 async fn check_connection(&self, context: &ConnectionContext<'_>) -> Verdict {
250 self.decide(context.client_ip).await
251 }
252
253 async fn check_identifiers(&self, context: &IdentifierContext<'_>) -> Verdict {
254 self.decide(context.client_ip).await
255 }
256}
257
258#[cfg(test)]
259mod tests {
260 use super::*;
261 use axum::http::Method;
262 use std::collections::HashMap;
263
264 #[derive(Default)]
266 struct StubResolver {
267 ptr: HashMap<IpAddr, Vec<String>>,
268 forward: HashMap<String, Vec<IpAddr>>,
269 ptr_error: Option<String>,
270 forward_error: Option<String>,
271 hang: bool,
272 }
273
274 impl StubResolver {
278 fn with_ptr(mut self, ip: &str, names: &[&str]) -> Self {
279 self.ptr.insert(
280 ip.parse().unwrap(),
281 names.iter().map(std::string::ToString::to_string).collect(),
282 );
283 self
284 }
285
286 fn with_forward(mut self, name: &str, ips: &[&str]) -> Self {
287 self.forward.insert(
288 name.to_string(),
289 ips.iter().map(|ip| ip.parse().unwrap()).collect(),
290 );
291 self
292 }
293 }
294
295 #[async_trait]
296 impl Resolver for StubResolver {
297 async fn reverse(&self, ip: IpAddr) -> Result<Vec<String>, String> {
298 if self.hang {
299 tokio::time::sleep(Duration::from_secs(3600)).await;
301 }
302 if let Some(error) = &self.ptr_error {
303 return Err(error.clone());
304 }
305 Ok(self.ptr.get(&ip).cloned().unwrap_or_default())
306 }
307
308 async fn forward(&self, name: &str) -> Result<Vec<IpAddr>, String> {
309 if let Some(error) = &self.forward_error {
310 return Err(error.clone());
311 }
312 Ok(self.forward.get(name).cloned().unwrap_or_default())
313 }
314
315 async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
316 unreachable!("the reverse_dns filter never looks up TXT records")
317 }
318 }
319
320 fn filter(cfg: &Settings, resolver: StubResolver) -> ClientHasValidReverseDns {
321 ClientHasValidReverseDns::with_resolver("ptr", cfg, Arc::new(resolver)).unwrap()
322 }
323
324 async fn check(filter: &ClientHasValidReverseDns, ip: &str) -> Verdict {
325 filter
326 .check_connection(&ConnectionContext {
327 client_ip: Some(ip.parse().unwrap()),
328 method: &Method::POST,
329 path: "/newOrder",
330 })
331 .await
332 }
333
334 fn assert_denied(verdict: Verdict, needle: &str) {
335 match verdict {
336 Verdict::Fail(detail) => {
337 assert!(detail.contains(needle), "{detail:?} lacks {needle:?}");
338 }
339 other => panic!("expected Fail, got {other:?}"),
340 }
341 }
342
343 fn assert_internal(verdict: Verdict, needle: &str) {
347 match verdict {
348 Verdict::Undecided(detail) => {
349 assert!(detail.contains(needle), "{detail:?} lacks {needle:?}");
350 }
351 other => panic!("expected Undecided, got {other:?}"),
352 }
353 }
354
355 #[tokio::test]
356 async fn accepts_a_forward_confirmed_client() {
357 let resolver = StubResolver::default()
358 .with_ptr("203.0.113.9", &["host.example.com"])
359 .with_forward("host.example.com", &["203.0.113.9"]);
360 let filter = filter(&Settings::default(), resolver);
361 assert_eq!(check(&filter, "203.0.113.9").await, Verdict::Pass);
362 }
363
364 #[tokio::test]
365 async fn denies_a_client_without_a_ptr_record() {
366 let filter = filter(&Settings::default(), StubResolver::default());
367 assert_denied(check(&filter, "203.0.113.9").await, "no PTR");
368 }
369
370 #[tokio::test]
371 async fn denies_when_the_forward_record_does_not_come_back() {
372 let resolver = StubResolver::default()
374 .with_ptr("203.0.113.9", &["trusted.example.com"])
375 .with_forward("trusted.example.com", &["198.51.100.4"]);
376 let filter = filter(&Settings::default(), resolver);
377 assert_denied(check(&filter, "203.0.113.9").await, "does not resolve back");
378 }
379
380 #[tokio::test]
381 async fn accepts_an_unconfirmed_ptr_when_confirmation_is_off() {
382 let resolver = StubResolver::default().with_ptr("203.0.113.9", &["host.example.com"]);
383 let cfg = Settings {
384 require_forward_confirm: false,
385 ..Settings::default()
386 };
387 assert_eq!(
388 check(&filter(&cfg, resolver), "203.0.113.9").await,
389 Verdict::Pass
390 );
391 }
392
393 #[tokio::test]
394 async fn applies_the_hostname_allow_list() {
395 let resolver = StubResolver::default()
396 .with_ptr("203.0.113.9", &["host.other.net"])
397 .with_forward("host.other.net", &["203.0.113.9"]);
398 let cfg = Settings {
399 allow_regex: vec![r".*\.corp\.example\.com".to_string()],
400 ..Settings::default()
401 };
402 assert_denied(
403 check(&filter(&cfg, resolver), "203.0.113.9").await,
404 "is not allowed",
405 );
406 }
407
408 #[tokio::test]
409 async fn applies_the_hostname_deny_list() {
410 let resolver = StubResolver::default()
411 .with_ptr("203.0.113.9", &["bad.example.com"])
412 .with_forward("bad.example.com", &["203.0.113.9"]);
413 let cfg = Settings {
414 deny_regex: vec![r"bad\..*".to_string()],
415 ..Settings::default()
416 };
417 assert_denied(
418 check(&filter(&cfg, resolver), "203.0.113.9").await,
419 "is denied",
420 );
421 }
422
423 #[tokio::test]
431 async fn a_denied_name_cannot_be_masked_by_a_second_ptr_record() {
432 let resolver = StubResolver::default()
433 .with_ptr("203.0.113.9", &["bad.example.com", "ok.example.com"])
434 .with_forward("bad.example.com", &["203.0.113.9"])
435 .with_forward("ok.example.com", &["203.0.113.9"]);
436 let cfg = Settings {
437 deny_regex: vec![r"bad\..*".to_string()],
438 ..Settings::default()
439 };
440 assert_denied(
441 check(&filter(&cfg, resolver), "203.0.113.9").await,
442 "bad.example.com",
443 );
444 }
445
446 #[tokio::test]
448 async fn a_denied_name_refuses_whatever_order_the_records_arrive_in() {
449 let resolver = StubResolver::default()
450 .with_ptr("203.0.113.9", &["ok.example.com", "bad.example.com"])
451 .with_forward("bad.example.com", &["203.0.113.9"])
452 .with_forward("ok.example.com", &["203.0.113.9"]);
453 let cfg = Settings {
454 deny_regex: vec![r"bad\..*".to_string()],
455 ..Settings::default()
456 };
457 assert_denied(
458 check(&filter(&cfg, resolver), "203.0.113.9").await,
459 "bad.example.com",
460 );
461 }
462
463 #[tokio::test]
464 async fn one_acceptable_name_among_several_is_enough() {
465 let resolver = StubResolver::default()
466 .with_ptr("203.0.113.9", &["stale.example.com", "host.example.com"])
467 .with_forward("stale.example.com", &["198.51.100.4"])
468 .with_forward("host.example.com", &["203.0.113.9"]);
469 let filter = filter(&Settings::default(), resolver);
470 assert_eq!(check(&filter, "203.0.113.9").await, Verdict::Pass);
471 }
472
473 #[tokio::test]
474 async fn a_ptr_lookup_failure_is_internal_not_a_denial() {
475 let resolver = StubResolver {
476 ptr_error: Some("SERVFAIL".to_string()),
477 ..StubResolver::default()
478 };
479 let filter = filter(&Settings::default(), resolver);
480 assert_internal(check(&filter, "203.0.113.9").await, "SERVFAIL");
481 }
482
483 #[tokio::test]
484 async fn a_forward_lookup_failure_is_internal() {
485 let resolver = StubResolver {
486 forward_error: Some("SERVFAIL".to_string()),
487 ..StubResolver::default().with_ptr("203.0.113.9", &["host.example.com"])
488 };
489 let filter = filter(&Settings::default(), resolver);
490 assert_internal(check(&filter, "203.0.113.9").await, "forward lookup failed");
491 }
492
493 #[tokio::test]
494 async fn a_wedged_resolver_times_out_rather_than_hanging() {
495 let resolver = StubResolver {
496 hang: true,
497 ..StubResolver::default()
498 };
499 let cfg = Settings {
500 timeout_ms: 10,
501 ..Settings::default()
502 };
503 assert_internal(
504 check(&filter(&cfg, resolver), "203.0.113.9").await,
505 "timed out",
506 );
507 }
508
509 #[tokio::test]
510 async fn a_missing_client_address_is_denied() {
511 let filter = filter(&Settings::default(), StubResolver::default());
512 let error = filter
513 .check_connection(&ConnectionContext {
514 client_ip: None,
515 method: &Method::POST,
516 path: "/newOrder",
517 })
518 .await;
519 assert_denied(error, "unavailable");
520 }
521
522 #[tokio::test]
523 async fn an_ipv4_mapped_client_is_canonicalized_before_lookup() {
524 let resolver = StubResolver::default()
525 .with_ptr("192.168.1.5", &["host.example.com"])
526 .with_forward("host.example.com", &["192.168.1.5"]);
527 let filter = filter(&Settings::default(), resolver);
528 assert_eq!(check(&filter, "::ffff:192.168.1.5").await, Verdict::Pass);
529 }
530
531 #[test]
532 fn a_bad_hostname_pattern_is_a_startup_error() {
533 let cfg = Settings {
534 deny_regex: vec!["[unclosed".to_string()],
535 ..Settings::default()
536 };
537 let error =
538 ClientHasValidReverseDns::with_resolver("ptr", &cfg, Arc::new(StubResolver::default()))
539 .unwrap_err()
540 .to_string();
541 assert!(error.contains("filter.check.ptr.deny_regex"), "{error}");
542 }
543
544 #[test]
545 fn reports_its_type_and_stages() {
546 let check = filter(&Settings::default(), StubResolver::default());
547 assert_eq!(check.kind(), "reverse_dns");
548 assert_eq!(check.stages(), StageSet::connection_only());
552 }
553
554 #[test]
557 fn the_filter_debug_shows_its_policy() {
558 let filter = ClientHasValidReverseDns::with_resolver(
559 "ptr",
560 &Settings {
561 require_forward_confirm: true,
562 allow_regex: vec![r"host\.example\.com".to_string()],
563 timeout_ms: 1234,
564 ..Settings::default()
565 },
566 Arc::new(StubResolver::default()),
567 )
568 .unwrap();
569
570 let rendered = format!("{filter:?}");
571 assert!(rendered.contains("ClientHasValidReverseDns"), "{rendered}");
572 assert!(rendered.contains("true"), "{rendered}");
573 assert!(rendered.contains("1.234s"), "{rendered}");
574 }
575}