Skip to main content

rsipstack/resolver/
sip_resolver.rs

1use crate::sip::{Domain, Port, Transport};
2use hickory_resolver::{
3    config::{LookupIpStrategy, NameServerConfig, ResolverConfig},
4    net::runtime::TokioRuntimeProvider,
5    proto::rr::RData,
6    TokioResolver,
7};
8use rand::RngExt;
9use std::net::IpAddr;
10use std::net::SocketAddr;
11use std::str::FromStr;
12use std::sync::Arc;
13
14#[derive(Debug, Clone, PartialEq, Eq, Hash)]
15pub struct Target {
16    pub addr: SocketAddr,
17    pub transport: Transport,
18}
19
20#[derive(Debug, Clone)]
21pub struct SipResolver {
22    backend: ResolverBackend,
23}
24
25#[derive(Debug, Clone)]
26enum ResolverBackend {
27    Hickory(Arc<TokioResolver>),
28    System,
29}
30
31impl Default for SipResolver {
32    fn default() -> Self {
33        Self::new()
34    }
35}
36
37impl SipResolver {
38    pub fn new() -> Self {
39        match Self::try_hickory() {
40            Ok(resolver) => resolver,
41            Err(e) => {
42                tracing::warn!(
43                    error = %e,
44                    "Failed to build hickory DNS resolver, falling back to system resolver"
45                );
46                Self::system()
47            }
48        }
49    }
50
51    /// Build a resolver that queries the given nameservers, bypassing the
52    /// system DNS configuration (which may contain unparseable entries such
53    /// as IPv6 link-local addresses with a `%zone` scope suffix).
54    pub fn with_nameservers(nameservers: Vec<IpAddr>) -> Self {
55        let mut config = ResolverConfig::default();
56        for ip in nameservers {
57            config.add_name_server(NameServerConfig::udp_and_tcp(ip));
58        }
59
60        let mut builder =
61            TokioResolver::builder_with_config(config, TokioRuntimeProvider::default());
62        builder.options_mut().ip_strategy = LookupIpStrategy::Ipv4thenIpv6;
63
64        match builder.build() {
65            Ok(resolver) => Self {
66                backend: ResolverBackend::Hickory(Arc::new(resolver)),
67            },
68            Err(e) => {
69                tracing::warn!(
70                    error = %e,
71                    "Failed to build hickory DNS resolver with custom nameservers, falling back to system resolver"
72                );
73                Self::system()
74            }
75        }
76    }
77
78    fn system() -> Self {
79        Self {
80            backend: ResolverBackend::System,
81        }
82    }
83
84    fn try_hickory() -> Result<Self, String> {
85        let mut builder = TokioResolver::builder_tokio().map_err(|e| e.to_string())?;
86        builder.options_mut().ip_strategy = LookupIpStrategy::Ipv4thenIpv6;
87
88        let resolver = builder.build().map_err(|e| e.to_string())?;
89        Ok(Self {
90            backend: ResolverBackend::Hickory(Arc::new(resolver)),
91        })
92    }
93
94    /// Main lookup function implementing core of RFC 3263 (SRV + Fallback)
95    pub async fn lookup(
96        &self,
97        domain: &Domain,
98        port: Option<Port>,
99        transport: Option<Transport>,
100        secure: bool,
101    ) -> Result<Vec<Target>, String> {
102        match &self.backend {
103            ResolverBackend::Hickory(resolver) => {
104                let source = HickorySource(resolver.clone());
105                resolve_logic(&source, domain, port, transport, secure).await
106            }
107            ResolverBackend::System => {
108                let source = SystemLookupSource;
109                resolve_logic(&source, domain, port, transport, secure).await
110            }
111        }
112    }
113}
114
115#[derive(Debug, Clone, Eq, PartialEq)]
116pub struct SrvRecord {
117    pub target: String,
118    pub port: u16,
119    pub priority: u16,
120    pub weight: u16,
121}
122
123#[async_trait::async_trait]
124pub trait LookupSource: Send + Sync {
125    async fn lookup_srv(&self, name: &str) -> Result<Vec<SrvRecord>, String>;
126    async fn lookup_a_aaaa(&self, name: &str) -> Result<Vec<IpAddr>, String>;
127}
128
129struct HickorySource(Arc<TokioResolver>);
130
131#[async_trait::async_trait]
132impl LookupSource for HickorySource {
133    async fn lookup_srv(&self, name: &str) -> Result<Vec<SrvRecord>, String> {
134        match self.0.srv_lookup(name).await {
135            Ok(records) => {
136                let mut res = Vec::new();
137                for r in records.message().all_sections() {
138                    if let RData::SRV(srv) = &r.data {
139                        let target = srv.target.to_string();
140                        // Remove trailing dot
141                        let target = target.trim_end_matches('.').to_string();
142                        res.push(SrvRecord {
143                            target,
144                            port: srv.port,
145                            priority: srv.priority,
146                            weight: srv.weight,
147                        });
148                    }
149                }
150                Ok(res)
151            }
152            Err(e) => Err(e.to_string()),
153        }
154    }
155
156    async fn lookup_a_aaaa(&self, name: &str) -> Result<Vec<IpAddr>, String> {
157        match self.0.lookup_ip(name).await {
158            Ok(records) => Ok(records.iter().collect()),
159            Err(e) => Err(e.to_string()),
160        }
161    }
162}
163
164/// Fallback lookup source backed by the OS resolver (`getaddrinfo` via
165/// `tokio::net::lookup_host`). Does not support SRV lookups; callers should
166/// fall back to A/AAAA resolution.
167struct SystemLookupSource;
168
169#[async_trait::async_trait]
170impl LookupSource for SystemLookupSource {
171    async fn lookup_srv(&self, name: &str) -> Result<Vec<SrvRecord>, String> {
172        Err(format!(
173            "SRV lookup not supported by system resolver: {}",
174            name
175        ))
176    }
177
178    async fn lookup_a_aaaa(&self, name: &str) -> Result<Vec<IpAddr>, String> {
179        let addr_str = format!("{}:0", name);
180        let result = tokio::net::lookup_host(&addr_str).await;
181        match result {
182            Ok(addrs) => Ok(addrs.map(|a| a.ip()).collect()),
183            Err(e) => Err(format!("DNS resolution failed for {}: {}", name, e)),
184        }
185    }
186}
187
188pub async fn resolve_logic<S: LookupSource + ?Sized>(
189    source: &S,
190    domain: &Domain,
191    port: Option<Port>,
192    transport: Option<Transport>,
193    secure: bool,
194) -> Result<Vec<Target>, String> {
195    let domain_str = domain.to_string();
196
197    if let Ok(ip) = IpAddr::from_str(&domain_str) {
198        let t = transport.unwrap_or(if secure {
199            Transport::Tls
200        } else {
201            Transport::Udp
202        });
203        let p: u16 = port
204            .map(|p| p.into())
205            .unwrap_or_else(|| t.default_port().into());
206        return Ok(vec![Target {
207            addr: SocketAddr::new(ip, p),
208            transport: t,
209        }]);
210    }
211
212    if let Some(p) = port {
213        let t = transport.unwrap_or(if secure {
214            Transport::Tls
215        } else {
216            Transport::Udp
217        });
218        let ips = source.lookup_a_aaaa(&domain_str).await.unwrap_or_default();
219
220        if ips.is_empty() {
221            return Err(format!("Could not resolve IP for {}", domain_str));
222        }
223
224        let p_u16: u16 = p.into();
225        let targets = ips
226            .into_iter()
227            .map(|ip| Target {
228                addr: SocketAddr::new(ip, p_u16),
229                transport: t,
230            })
231            .collect();
232        return Ok(targets);
233    }
234
235    let mut targets = Vec::new();
236    let mut candidates = Vec::new();
237
238    if let Some(t) = transport {
239        candidates.push(t);
240    } else {
241        if secure {
242            candidates.push(Transport::Tls);
243        } else {
244            candidates.push(Transport::Udp);
245            candidates.push(Transport::Tcp);
246        }
247    }
248
249    let mut _srv_found = false;
250
251    for t in candidates.iter() {
252        let prefix = srv_prefix(*t, secure);
253        if prefix.is_empty() {
254            continue;
255        } // Unsupported transport for SRV
256
257        let srv_name = format!("{}.{}", prefix, domain_str);
258
259        if let Ok(records) = source.lookup_srv(&srv_name).await {
260            if !records.is_empty() {
261                _srv_found = true;
262                let ordered = order_srv_records(records);
263
264                for rec in ordered {
265                    // Start sub-query for A/AAAA
266                    if let Ok(ips) = source.lookup_a_aaaa(&rec.target).await {
267                        for ip in ips {
268                            targets.push(Target {
269                                addr: SocketAddr::new(ip, rec.port),
270                                transport: *t,
271                            });
272                        }
273                    }
274                }
275            }
276        }
277    }
278
279    if targets.is_empty() {
280        let def_transport = transport.unwrap_or(if secure {
281            Transport::Tls
282        } else {
283            Transport::Udp
284        });
285        let def_port = def_transport.default_port();
286
287        match source.lookup_a_aaaa(&domain_str).await {
288            Ok(ips) if !ips.is_empty() => {
289                for ip in ips {
290                    targets.push(Target {
291                        addr: SocketAddr::new(ip, def_port.into()),
292                        transport: def_transport,
293                    });
294                }
295                Ok(targets)
296            }
297            _ => Err(format!("Resolution failed for {}", domain_str)),
298        }
299    } else {
300        Ok(targets)
301    }
302}
303
304fn srv_prefix(transport: Transport, secure: bool) -> &'static str {
305    match (transport, secure) {
306        (Transport::Udp, false) => "_sip._udp",
307        (Transport::Tcp, false) => "_sip._tcp",
308        (Transport::Tls, _) => "_sips._tcp",
309        (Transport::Tcp, true) => "_sips._tcp",
310        (Transport::Wss, true) => "_sips._tcp", // Common practice fallback
311        _ => "",
312    }
313}
314
315fn order_srv_records(mut records: Vec<SrvRecord>) -> Vec<SrvRecord> {
316    records.sort_by_key(|k| k.priority);
317
318    let mut ordered = Vec::new();
319    let mut start_idx = 0;
320
321    while start_idx < records.len() {
322        let current_priority = records[start_idx].priority;
323        let mut end_idx = start_idx;
324
325        // Find range with same priority
326        while end_idx < records.len() && records[end_idx].priority == current_priority {
327            end_idx += 1;
328        }
329
330        // Group of records with same priority
331        let mut group = records[start_idx..end_idx].to_vec();
332
333        // Selection sort based on weights
334        while !group.is_empty() {
335            let total_weight: u32 = group.iter().map(|r| r.weight as u32).sum();
336            let mut rng = rand::rng();
337
338            if total_weight == 0 {
339                // All zero, just pick one (shuffle or first)
340                let idx = rng.random_range(0..group.len()); // 0..len (exclusive) => OK
341                ordered.push(group.remove(idx));
342            } else {
343                let mut r = rng.random_range(0..=total_weight); // 0..=total
344                let mut selected_idx = 0;
345                for (i, rec) in group.iter().enumerate() {
346                    let w = rec.weight as u32;
347                    if r <= w {
348                        selected_idx = i;
349                        break;
350                    }
351                    r -= w;
352                }
353                if selected_idx >= group.len() {
354                    selected_idx = group.len() - 1;
355                }
356
357                ordered.push(group.remove(selected_idx));
358            }
359        }
360
361        start_idx = end_idx;
362    }
363
364    ordered
365}
366
367#[cfg(test)]
368mod tests {
369    use super::*;
370    use parking_lot::Mutex;
371    use std::collections::HashMap;
372
373    struct MockDns {
374        srv: Mutex<HashMap<String, Vec<SrvRecord>>>,
375        a: Mutex<HashMap<String, Vec<IpAddr>>>,
376    }
377
378    impl MockDns {
379        fn new() -> Self {
380            Self {
381                srv: Mutex::new(HashMap::new()),
382                a: Mutex::new(HashMap::new()),
383            }
384        }
385
386        fn add_srv(&self, name: &str, target: &str, port: u16, priority: u16, weight: u16) {
387            let mut map = self.srv.lock();
388            map.entry(name.to_string()).or_default().push(SrvRecord {
389                target: target.to_string(),
390                port,
391                priority,
392                weight,
393            });
394        }
395
396        fn add_a(&self, name: &str, ip: IpAddr) {
397            let mut map = self.a.lock();
398            map.entry(name.to_string()).or_default().push(ip);
399        }
400    }
401
402    #[async_trait::async_trait]
403    impl LookupSource for MockDns {
404        async fn lookup_srv(&self, name: &str) -> Result<Vec<SrvRecord>, String> {
405            let map = self.srv.lock();
406            if let Some(recs) = map.get(name) {
407                Ok(recs.clone())
408            } else {
409                Err("Not found".to_string())
410            }
411        }
412
413        async fn lookup_a_aaaa(&self, name: &str) -> Result<Vec<IpAddr>, String> {
414            let map = self.a.lock();
415            if let Some(ips) = map.get(name) {
416                Ok(ips.clone())
417            } else {
418                Err("Not found".to_string())
419            }
420        }
421    }
422
423    #[tokio::test]
424    async fn test_ip_direct() {
425        let mock = MockDns::new();
426        let domain = Domain::from("127.0.0.1".to_string());
427
428        let res = resolve_logic(&mock, &domain, None, None, false)
429            .await
430            .unwrap();
431        assert_eq!(res.len(), 1);
432        assert_eq!(res[0].addr.ip().to_string(), "127.0.0.1");
433        assert_eq!(res[0].transport, Transport::Udp); // Default insecure
434    }
435
436    #[tokio::test]
437    async fn test_domain_with_port() {
438        let mock = MockDns::new();
439        mock.add_a("example.com", "1.2.3.4".parse().unwrap());
440
441        let domain = Domain::from("example.com".to_string());
442        let res = resolve_logic(
443            &mock,
444            &domain,
445            Some(5090.into()),
446            Some(Transport::Tcp),
447            false,
448        )
449        .await
450        .unwrap();
451
452        assert_eq!(res.len(), 1);
453        assert_eq!(res[0].addr, "1.2.3.4:5090".parse().unwrap());
454        assert_eq!(res[0].transport, Transport::Tcp);
455    }
456
457    #[tokio::test]
458    async fn test_srv_lookup_basic() {
459        let mock = MockDns::new();
460        // Setup SRV
461        mock.add_srv("_sip._udp.example.com", "sip1.example.com", 5060, 10, 100);
462
463        // Setup A
464        mock.add_a("sip1.example.com", "10.0.0.1".parse().unwrap());
465
466        let domain = Domain::from("example.com".to_string());
467        let res = resolve_logic(&mock, &domain, None, Some(Transport::Udp), false)
468            .await
469            .unwrap();
470
471        assert_eq!(res.len(), 1);
472        assert_eq!(res[0].addr, "10.0.0.1:5060".parse().unwrap());
473        assert_eq!(res[0].transport, Transport::Udp);
474    }
475
476    #[tokio::test]
477    async fn test_srv_priority() {
478        let mock = MockDns::new();
479        // Priority 10 vs 20
480        mock.add_srv("_sip._udp.example.com", "high.example.com", 5060, 10, 100);
481        mock.add_srv("_sip._udp.example.com", "low.example.com", 5060, 20, 100);
482
483        mock.add_a("high.example.com", "1.1.1.1".parse().unwrap());
484        mock.add_a("low.example.com", "2.2.2.2".parse().unwrap());
485
486        let domain = Domain::from("example.com".to_string());
487        let res = resolve_logic(&mock, &domain, None, Some(Transport::Udp), false)
488            .await
489            .unwrap();
490
491        assert_eq!(res.len(), 2);
492        assert_eq!(res[0].addr.ip().to_string(), "1.1.1.1");
493        assert_eq!(res[1].addr.ip().to_string(), "2.2.2.2");
494
495        // Check order
496        let ips: Vec<String> = res.iter().map(|t| t.addr.ip().to_string()).collect();
497        assert_eq!(ips, vec!["1.1.1.1", "2.2.2.2"]);
498    }
499
500    #[tokio::test]
501    async fn test_fallback_to_a() {
502        let mock = MockDns::new();
503        // No SRV records added
504        mock.add_a("example.com", "9.9.9.9".parse().unwrap());
505
506        let domain = Domain::from("example.com".to_string());
507        let res = resolve_logic(&mock, &domain, None, Some(Transport::Udp), false)
508            .await
509            .unwrap();
510
511        assert_eq!(res.len(), 1);
512        assert_eq!(res[0].addr, "9.9.9.9:5060".parse().unwrap());
513    }
514
515    #[test]
516    fn test_srv_ordering_weight() {
517        let records = vec![
518            SrvRecord {
519                target: "a".into(),
520                port: 1,
521                priority: 1,
522                weight: 10,
523            },
524            SrvRecord {
525                target: "b".into(),
526                port: 1,
527                priority: 1,
528                weight: 90,
529            },
530        ];
531
532        // This is randomized, but checking it runs without panic
533        let ordered = order_srv_records(records);
534        assert_eq!(ordered.len(), 2);
535    }
536
537    #[tokio::test]
538    async fn test_system_lookup_source_resolves_host() {
539        let source = SystemLookupSource;
540        let ips = source.lookup_a_aaaa("localhost").await.unwrap();
541        assert!(!ips.is_empty());
542    }
543
544    #[tokio::test]
545    async fn test_system_backend_resolves_ip() {
546        let resolver = SipResolver {
547            backend: ResolverBackend::System,
548        };
549        let domain = Domain::from("127.0.0.1".to_string());
550        let res = resolver.lookup(&domain, None, None, false).await.unwrap();
551        assert_eq!(res.len(), 1);
552        assert_eq!(res[0].addr.ip().to_string(), "127.0.0.1");
553        assert_eq!(res[0].transport, Transport::Udp);
554    }
555
556    #[test]
557    fn test_with_nameservers_builds() {
558        let resolver = SipResolver::with_nameservers(vec!["127.0.0.1".parse().unwrap()]);
559        assert!(matches!(resolver.backend, ResolverBackend::Hickory(_)));
560    }
561
562    #[test]
563    fn test_new_does_not_panic() {
564        let _ = SipResolver::new();
565    }
566}