tokio-xmpp: add option to set custom hickory resolver

This commit is contained in:
famfo 2026-05-21 13:57:26 +02:00 committed by pep
commit 559159d458
3 changed files with 64 additions and 18 deletions

View file

@ -12,6 +12,8 @@ Version NEXT:
- Expose `client_auth` method to allow manual stream setups for advanced - Expose `client_auth` method to allow manual stream setups for advanced
use cases. use cases.
- Add feature flag to disable validation in xmpp_parsers - Add feature flag to disable validation in xmpp_parsers
- Added option to set custom DNS resolver to
`tokio_xmpp::connect::DnsConfig`
* Fixed: * Fixed:
- Ignore missing "version" stream attribute for 0114 components. - Ignore missing "version" stream attribute for 0114 components.
- Gate `AsRawFd` behind `ktls` feature to make Windows build work again. - Gate `AsRawFd` behind `ktls` feature to make Windows build work again.

View file

@ -66,6 +66,7 @@ async fn main() {
host: jid.domain().as_str().to_owned(), host: jid.domain().as_str().to_owned(),
srv: "_xmpp-client._tcp".to_owned(), srv: "_xmpp-client._tcp".to_owned(),
fallback_port: 5222, fallback_port: 5222,
resolver: None,
}), }),
jid.clone().into(), jid.clone().into(),
password.clone(), password.clone(),

View file

@ -23,6 +23,8 @@ pub enum DnsConfig {
srv: String, srv: String,
/// When SRV resolution fails what port to use /// When SRV resolution fails what port to use
fallback_port: u16, fallback_port: u16,
/// Pre-configured DNS resolver
resolver: Option<TokioResolver>,
}, },
/// Manually define server host and port /// Manually define server host and port
@ -33,6 +35,8 @@ pub enum DnsConfig {
host: String, host: String,
/// Server port /// Server port
port: u16, port: u16,
/// Pre-configured DNS resolver
resolver: Option<TokioResolver>,
}, },
/// Manually define IP: port (TODO: socket) /// Manually define IP: port (TODO: socket)
@ -49,7 +53,7 @@ impl fmt::Display for DnsConfig {
#[cfg(feature = "dns")] #[cfg(feature = "dns")]
Self::UseSrv { host, .. } => write!(f, "{}", host), Self::UseSrv { host, .. } => write!(f, "{}", host),
#[cfg(feature = "dns")] #[cfg(feature = "dns")]
Self::NoSrv { host, port } => write!(f, "{}:{}", host, port), Self::NoSrv { host, port, .. } => write!(f, "{}:{}", host, port),
Self::Addr { addr } => write!(f, "{}", addr), Self::Addr { addr } => write!(f, "{}", addr),
} }
} }
@ -63,6 +67,7 @@ impl DnsConfig {
host: host.to_string(), host: host.to_string(),
srv: srv.to_string(), srv: srv.to_string(),
fallback_port, fallback_port,
resolver: None,
} }
} }
@ -73,6 +78,7 @@ impl DnsConfig {
host: host.to_string(), host: host.to_string(),
srv: "_xmpp-client._tcp".to_string(), srv: "_xmpp-client._tcp".to_string(),
fallback_port: 5222, fallback_port: 5222,
resolver: None,
} }
} }
@ -83,6 +89,7 @@ impl DnsConfig {
host: host.to_string(), host: host.to_string(),
srv: "_xmpps-client._tcp".to_string(), srv: "_xmpps-client._tcp".to_string(),
fallback_port: 5223, fallback_port: 5223,
resolver: None,
} }
} }
@ -92,6 +99,7 @@ impl DnsConfig {
Self::NoSrv { Self::NoSrv {
host: host.to_string(), host: host.to_string(),
port, port,
resolver: None,
} }
} }
@ -102,6 +110,20 @@ impl DnsConfig {
} }
} }
/// Set pre-configured DNS resolver
#[cfg(feature = "dns")]
pub fn with_resolver(&mut self, custom_resolver: TokioResolver) {
match self {
Self::UseSrv {
ref mut resolver, ..
} => *resolver = Some(custom_resolver),
Self::NoSrv {
ref mut resolver, ..
} => *resolver = Some(custom_resolver),
Self::Addr { .. } => {}
}
}
/// Try resolve the DnsConfig to a TcpStream /// Try resolve the DnsConfig to a TcpStream
pub async fn resolve(&self) -> Result<TcpStream, Error> { pub async fn resolve(&self) -> Result<TcpStream, Error> {
match self { match self {
@ -110,9 +132,14 @@ impl DnsConfig {
host, host,
srv, srv,
fallback_port, fallback_port,
} => Self::resolve_srv(host, srv, *fallback_port).await, resolver,
} => Self::resolve_srv(host, srv, *fallback_port, resolver).await,
#[cfg(feature = "dns")] #[cfg(feature = "dns")]
Self::NoSrv { host, port } => Self::resolve_no_srv(host, *port).await, Self::NoSrv {
host,
port,
resolver,
} => Self::resolve_no_srv(host, *port, resolver).await,
Self::Addr { addr } => { Self::Addr { addr } => {
// TODO: Unix domain socket // TODO: Unix domain socket
let addr: SocketAddr = addr.parse()?; let addr: SocketAddr = addr.parse()?;
@ -122,7 +149,12 @@ impl DnsConfig {
} }
#[cfg(feature = "dns")] #[cfg(feature = "dns")]
async fn resolve_srv(host: &str, srv: &str, fallback_port: u16) -> Result<TcpStream, Error> { async fn resolve_srv(
host: &str,
srv: &str,
fallback_port: u16,
resolver: &Option<TokioResolver>,
) -> Result<TcpStream, Error> {
let ascii_domain = idna::domain_to_ascii(host)?; let ascii_domain = idna::domain_to_ascii(host)?;
if let Ok(ip) = ascii_domain.parse() { if let Ok(ip) = ascii_domain.parse() {
@ -130,14 +162,11 @@ impl DnsConfig {
return Ok(TcpStream::connect(&SocketAddr::new(ip, fallback_port)).await?); return Ok(TcpStream::connect(&SocketAddr::new(ip, fallback_port)).await?);
} }
let (_config, options) = hickory_resolver::system_conf::read_system_conf()?; let resolver = Self::new_resolver(resolver)?;
let resolver = TokioResolver::builder_tokio()?
.with_options(options)
.build()?;
let srv_domain = format!("{}.{}.", srv, ascii_domain).into_name()?; let srv_domain = format!("{}.{}.", srv, ascii_domain).into_name()?;
let srv_records = resolver.srv_lookup(srv_domain.clone()).await.ok(); let srv_records = resolver.srv_lookup(srv_domain.clone()).await.ok();
let resolver_ref = Some(resolver);
match srv_records { match srv_records {
Some(lookup) => { Some(lookup) => {
// TODO: sort lookup records by priority/weight // TODO: sort lookup records by priority/weight
@ -147,7 +176,8 @@ impl DnsConfig {
}; };
debug!("Attempting connection to {srv_domain} {srv}"); debug!("Attempting connection to {srv_domain} {srv}");
if let Ok(stream) = Self::resolve_no_srv(&srv.target.to_ascii(), srv.port).await if let Ok(stream) =
Self::resolve_no_srv(&srv.target.to_ascii(), srv.port, &resolver_ref).await
{ {
return Ok(stream); return Ok(stream);
} }
@ -157,25 +187,24 @@ impl DnsConfig {
None => { None => {
// SRV lookup error, retry with hostname // SRV lookup error, retry with hostname
debug!("Attempting connection to {host}:{fallback_port}"); debug!("Attempting connection to {host}:{fallback_port}");
Self::resolve_no_srv(host, fallback_port).await Self::resolve_no_srv(host, fallback_port, &resolver_ref).await
} }
} }
} }
#[cfg(feature = "dns")] #[cfg(feature = "dns")]
async fn resolve_no_srv(host: &str, port: u16) -> Result<TcpStream, Error> { async fn resolve_no_srv(
host: &str,
port: u16,
resolver: &Option<TokioResolver>,
) -> Result<TcpStream, Error> {
let ascii_domain = idna::domain_to_ascii(host)?; let ascii_domain = idna::domain_to_ascii(host)?;
if let Ok(ip) = ascii_domain.parse() { if let Ok(ip) = ascii_domain.parse() {
return Ok(TcpStream::connect(&SocketAddr::new(ip, port)).await?); return Ok(TcpStream::connect(&SocketAddr::new(ip, port)).await?);
} }
let (_config, mut options) = hickory_resolver::system_conf::read_system_conf()?; let resolver = Self::new_resolver(resolver)?;
options.ip_strategy = LookupIpStrategy::Ipv4AndIpv6;
let resolver = TokioResolver::builder_tokio()?
.with_options(options)
.build()?;
let ips = resolver.lookup_ip(ascii_domain).await?; let ips = resolver.lookup_ip(ascii_domain).await?;
// Happy Eyeballs: connect to all records in parallel, return the // Happy Eyeballs: connect to all records in parallel, return the
@ -188,4 +217,18 @@ impl DnsConfig {
.map(|(result, _)| result) .map(|(result, _)| result)
.map_err(|_| Error::Disconnected) .map_err(|_| Error::Disconnected)
} }
#[cfg(feature = "dns")]
fn new_resolver(resolver: &Option<TokioResolver>) -> Result<TokioResolver, Error> {
if let Some(resolver) = resolver {
return Ok(resolver.clone());
}
let (_config, mut options) = hickory_resolver::system_conf::read_system_conf()?;
options.ip_strategy = LookupIpStrategy::Ipv4AndIpv6;
Ok(TokioResolver::builder_tokio()?
.with_options(options)
.build()?)
}
} }