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 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 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 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
164struct 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 } 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 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", _ => "",
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 while end_idx < records.len() && records[end_idx].priority == current_priority {
327 end_idx += 1;
328 }
329
330 let mut group = records[start_idx..end_idx].to_vec();
332
333 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 let idx = rng.random_range(0..group.len()); ordered.push(group.remove(idx));
342 } else {
343 let mut r = rng.random_range(0..=total_weight); 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); }
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 mock.add_srv("_sip._udp.example.com", "sip1.example.com", 5060, 10, 100);
462
463 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 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 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 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 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}