use anyhow::{bail, Context, Result}; use hickory_resolver::net::runtime::TokioRuntimeProvider; use hickory_resolver::proto::rr::rdata::SRV; use hickory_resolver::TokioResolver; use reqwest::Client; use std::{fmt, net::IpAddr}; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, net::TcpStream, time::{timeout, Duration}, }; use tracing::{debug, info, warn}; use tracing_subscriber::filter::LevelFilter; const DEFAULT_TS3_PROTO: &str = "udp"; const DEFAULT_TSDNS_PROTO: &str = "tcp"; const NICK_RESOLVE_URL: &str = "https://named.myteamspeak.com/lookup"; const TSDNS_PORT: u16 = 41144; const TSDNS_TERMINATOR: &[u8] = b"\n\r\r\r\n"; const TSDNS_TIMEOUT: Duration = Duration::from_secs(3); const TS3_SRV_FALLBACK_DISCOVERY_BUDGET: Duration = Duration::from_millis(900); const TSDNS_FALLBACK_DISCOVERY_BUDGET: Duration = Duration::from_millis(900); pub const DEFAULT_TEAMSPEAK_PORT: u16 = 9987; #[derive(Debug, Clone, PartialEq, Eq)] pub struct Args { pub host: String, pub service: String, pub protocol: String, } #[derive(Debug, Clone)] pub struct BuildInfo { pub version: String, pub commit: String, pub date: String, } #[derive(Debug, Clone, PartialEq, Eq)] pub struct SrvRecord { pub priority: u16, pub weight: u16, pub port: u16, pub target: String, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ClientResolutionMethod { Direct, Ts3Srv, TsdnsSrv, TsdnsTcp, Dns, Nick, } #[derive(Debug, Clone, PartialEq, Eq)] pub struct ClientResolution { pub input: String, pub address: String, pub method: ClientResolutionMethod, pub resolution: Option, } impl SrvRecord { pub fn resolved(&self) -> String { format_host_port(self.target.trim_end_matches('.'), self.port) } } impl Resolution { pub fn connection_address(&self) -> String { match self { Resolution::Dns { selected, .. } => { format_host_port(&selected.to_string(), DEFAULT_TEAMSPEAK_PORT) } Resolution::Srv { selected, .. } | Resolution::Nick { selected, .. } => { selected.resolved() } } } } #[derive(Debug, Clone, PartialEq, Eq)] pub enum Resolution { Dns { host: String, addresses: Vec, selected: IpAddr, }, Srv { service: String, host: String, protocol: String, records: Vec, selected: SrvRecord, }, Nick { nick: String, address: String, protocol: String, records: Vec, selected: SrvRecord, }, } #[derive(Debug, Clone)] pub struct ChanoraResolver { resolver: Option, http: Client, } impl fmt::Display for BuildInfo { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!( f, "{{version: {}, commit: {}, date: {}}}", self.version, self.commit, self.date ) } } pub fn build_info() -> BuildInfo { BuildInfo { version: option_env!("CHANORA_RESOLVER_BUILD_VERSION") .unwrap_or(env!("CARGO_PKG_VERSION")) .to_string(), commit: option_env!("CHANORA_RESOLVER_BUILD_COMMIT") .unwrap_or("unknown") .to_string(), date: option_env!("CHANORA_RESOLVER_BUILD_DATE") .unwrap_or("unknown") .to_string(), } } pub fn setup_log() { let level = std::env::var("CHANORA_RESOLVER_LOG_LEVEL").unwrap_or_else(|_| "info".to_string()); let filter = match level.to_lowercase().as_str() { "debug" => LevelFilter::DEBUG, "info" => LevelFilter::INFO, "warn" => LevelFilter::WARN, "error" => LevelFilter::ERROR, _ => LevelFilter::INFO, }; let _ = tracing_subscriber::fmt() .with_max_level(filter) .with_ansi(false) .with_target(false) .without_time() .with_writer(std::io::stdout) .try_init(); } impl ChanoraResolver { pub fn new() -> Result { let resolver = match srv_resolver_builder() { Ok(builder) => match builder.build() { Ok(resolver) => Some(resolver), Err(err) => { warn!(err = %err, "srv resolver unavailable"); None } }, Err(err) => { warn!(err = %err, "srv resolver unavailable"); None } }; Ok(Self { resolver, http: Client::new(), }) } pub async fn resolve(&self, args: &Args) -> Result { let normalized = normalize_args(args.clone()); validate_args(&normalized)?; debug!(build = %build_info(), "app start"); debug!( host = %normalized.host, service = %normalized.service, protocol = %normalized.protocol, "normalized flags" ); self.resolve_normalized(&normalized).await } pub async fn resolve_connection_address(&self, args: &Args) -> Result { Ok(self.resolve(args).await?.connection_address()) } pub async fn resolve_client_address(&self, input: &str) -> Result { Ok(self.resolve_client_request(input).await?.address) } pub async fn resolve_client_request(&self, input: &str) -> Result { let request = normalize_client_input(input)?; let host = request.host; info!(input = %input, host = %host, "client request resolution started"); if request.port.is_none() && should_try_server_name(&host) { match self.lookup_server_name(&host).await { Ok(Some(alias)) => { info!( input = %input, alias_host = %alias.host, alias_port = ?alias.port, "server name resolved" ); return self .resolve_client_host( input, alias.host, match alias.port { PortResolution::Explicit(port) => Some(port), PortResolution::KeepInput => None, }, true, ) .await; } Ok(None) if is_bare_numeric_name(&host) => { bail!("could not resolve TeamSpeak server name '{host}': not found"); } Ok(None) => {} Err(err) if is_bare_numeric_name(&host) => { bail!("could not resolve TeamSpeak server name '{host}': {err}"); } Err(err) => { debug!(err = %err, "server name lookup failed, continuing with host resolution"); } } } self.resolve_client_host(input, host, request.port, false) .await } async fn resolve_client_host( &self, input: &str, host: String, port: Option, from_server_name: bool, ) -> Result { let mut errors = Vec::new(); let fallback_port = port.unwrap_or(DEFAULT_TEAMSPEAK_PORT); if host.parse::().is_ok() { return Ok(client_resolution_with_optional_resolution( input, if from_server_name { ClientResolutionMethod::Nick } else { ClientResolutionMethod::Direct }, format_host_port(&host, fallback_port), None, )); } let mut direct = if is_bare_numeric_name(&host) { errors.push("dns: skipped bare numeric server name".to_string()); None } else { match self.resolve_dns(&host).await { Ok(resolution) => Some(resolution), Err(err) => { errors.push(format!("dns: {err}")); None } } }; if should_return_direct_dns_before_discovery(port) { if let Some(resolution) = direct.take() { return Ok(direct_client_resolution( input, from_server_name, fallback_port, resolution, )); } } let ts3 = self.resolve_ts3(&host, ""); match ts3_srv_discovery_budget(direct.is_some()) { Some(budget) => match timeout(budget, ts3).await { Ok(Ok(resolution)) => { let selected = match &resolution { Resolution::Srv { selected, .. } => selected.clone(), _ => unreachable!("ts3 resolution must be an srv result"), }; return self .client_srv_resolution( input, ClientResolutionMethod::Ts3Srv, &selected, resolution, ) .await; } Ok(Err(err)) => errors.push(format!("ts3 srv: {err}")), Err(_) => errors.push(format!( "ts3 srv: discovery timed out after {}ms", budget.as_millis() )), }, None => match ts3.await { Ok(resolution) => { let selected = match &resolution { Resolution::Srv { selected, .. } => selected.clone(), _ => unreachable!("ts3 resolution must be an srv result"), }; return self .client_srv_resolution( input, ClientResolutionMethod::Ts3Srv, &selected, resolution, ) .await; } Err(err) => errors.push(format!("ts3 srv: {err}")), }, } let tsdns = self.resolve_tsdns_candidates(input, &host, fallback_port); match tsdns_discovery_budget(direct.is_some()) { Some(budget) => match timeout(budget, tsdns).await { Ok(Ok(resolution)) => return Ok(resolution), Ok(Err(err)) => errors.push(format!("tsdns: {err}")), Err(_) => errors.push(format!( "tsdns: discovery timed out after {}ms", budget.as_millis() )), }, None => match tsdns.await { Ok(resolution) => return Ok(resolution), Err(err) => errors.push(format!("tsdns: {err}")), }, } if let Some(resolution) = direct { return Ok(direct_client_resolution( input, from_server_name, fallback_port, resolution, )); } bail!( "could not resolve TeamSpeak address '{host}': {}", errors.join("; ") ) } async fn resolve_tsdns_candidates( &self, input: &str, query_host: &str, fallback_port: u16, ) -> Result { let mut errors = Vec::new(); match self .resolve_tsdns_srv_candidates(input, query_host, fallback_port) .await { Ok(resolution) => return Ok(resolution), Err(err) => errors.push(format!("tsdns srv: {err}")), } match self .resolve_tsdns_tcp_candidates(input, query_host, fallback_port) .await { Ok(resolution) => return Ok(resolution), Err(err) => errors.push(format!("tsdns tcp: {err}")), } bail!("{}", errors.join("; ")) } async fn client_srv_resolution( &self, input: &str, method: ClientResolutionMethod, selected: &SrvRecord, resolution: Resolution, ) -> Result { let target = selected.target.trim_end_matches('.'); let addresses = self .lookup_ips(target, selected.port) .await .with_context(|| format!("srv target dns lookup failed for {target}"))?; if addresses.is_empty() { bail!("no addresses found for srv target: {target}"); } Ok(client_resolution_with_address( input, method, format_host_port(&addresses[0].to_string(), selected.port), resolution, )) } async fn resolve_tsdns_srv_candidates( &self, input: &str, query_host: &str, fallback_port: u16, ) -> Result { let mut errors = Vec::new(); for srv_host in tsdns_candidate_hosts(query_host) { match self.resolve_tsdns(&srv_host, "").await { Ok(resolution) => { return self .client_tsdns_srv_resolution(input, query_host, fallback_port, resolution) .await; } Err(err) => errors.push(format!("{srv_host}: {err}")), } } bail!("{}", errors.join("; ")) } async fn client_tsdns_srv_resolution( &self, input: &str, query_host: &str, fallback_port: u16, resolution: Resolution, ) -> Result { let selected = match &resolution { Resolution::Srv { selected, .. } => selected, _ => bail!("tsdns srv resolution must be an srv result"), }; let server = selected.target.trim_end_matches('.'); let endpoint = self .query_tsdns_server(server, selected.port, query_host) .await?; self.client_resolved_tsdns_endpoint(input, endpoint, fallback_port, Some(resolution)) .await } async fn resolve_tsdns_tcp_candidates( &self, input: &str, query_host: &str, fallback_port: u16, ) -> Result { let mut errors = Vec::new(); for server in tsdns_candidate_hosts(query_host) { match self .query_tsdns_server(&server, TSDNS_PORT, query_host) .await { Ok(endpoint) => { return self .client_resolved_tsdns_endpoint(input, endpoint, fallback_port, None) .await; } Err(err) => errors.push(format!("{server}: {err}")), } } bail!("{}", errors.join("; ")) } async fn client_resolved_tsdns_endpoint( &self, input: &str, endpoint: ServerNameResolution, fallback_port: u16, resolution: Option, ) -> Result { let method = if resolution.is_some() { ClientResolutionMethod::TsdnsSrv } else { ClientResolutionMethod::TsdnsTcp }; let port = endpoint.port.resolve(fallback_port); let address = if endpoint.host.parse::().is_ok() { format_host_port(&endpoint.host, port) } else { let addresses = self .lookup_ips(&endpoint.host, port) .await .with_context(|| format!("dns lookup failed for tsdns result {}", endpoint.host))?; if addresses.is_empty() { bail!("no addresses found for tsdns result: {}", endpoint.host); } format_host_port(&addresses[0].to_string(), port) }; Ok(client_resolution_with_optional_resolution( input, method, address, resolution, )) } async fn query_tsdns_server( &self, server: &str, port: u16, query_host: &str, ) -> Result { let addresses = self.lookup_ips(server, port).await?; let mut errors = Vec::new(); for addr in addresses { let socket = std::net::SocketAddr::new(addr, port); match self.query_tsdns_socket(socket, query_host).await { Ok(endpoint) => return Ok(endpoint), Err(err) => errors.push(format!("{socket}: {err}")), } } bail!("{}", errors.join("; ")) } async fn query_tsdns_socket( &self, socket: std::net::SocketAddr, query_host: &str, ) -> Result { let mut stream = timeout(TSDNS_TIMEOUT, TcpStream::connect(socket)) .await .context("tsdns connect timed out")? .context("tsdns connect failed")?; let mut request = query_host.trim().to_lowercase().into_bytes(); request.extend_from_slice(TSDNS_TERMINATOR); timeout(TSDNS_TIMEOUT, stream.write_all(&request)) .await .context("tsdns write timed out")? .context("tsdns write failed")?; let mut buffer = vec![0; 512]; let size = timeout(TSDNS_TIMEOUT, stream.read(&mut buffer)) .await .context("tsdns read timed out")? .context("tsdns read failed")?; buffer.truncate(size); let value = String::from_utf8(buffer).context("tsdns response returned invalid utf-8")?; let value = value.trim(); if value.is_empty() { bail!("empty tsdns response"); } if value == "404" { bail!("tsdns name not found"); } parse_resolved_server_name(value) } async fn lookup_server_name(&self, name: &str) -> Result> { if name.is_empty() { bail!("host cannot be empty"); } info!(name = %name, "trying to resolve server name"); let value = lookup_named_value(&self.http, name, true, "server name") .await? .ok_or_else(|| anyhow::anyhow!("server name lookup returned no data"))?; let endpoint = parse_resolved_server_name(value.trim())?; info!( name = %name, host = %endpoint.host, port = ?endpoint.port, "server name successfully resolved" ); Ok(Some(endpoint)) } async fn resolve_normalized(&self, args: &Args) -> Result { match args.service.as_str() { "dns" => self.resolve_dns(&args.host).await, "ts3" => self.resolve_ts3(&args.host, &args.protocol).await, "tsdns" => self.resolve_tsdns(&args.host, &args.protocol).await, "nick" => self.resolve_nick(&args.host, &args.protocol).await, other => bail!("unsupported service: {other}"), } } pub async fn resolve_dns(&self, host: &str) -> Result { info!(host = %host, service = "dns", "dns lookup started"); let lookup = tokio::net::lookup_host((host, DEFAULT_TEAMSPEAK_PORT)) .await .context("dns lookup failed")?; let addrs = collect_ips(lookup); if addrs.is_empty() { bail!("no addresses found for host: {host}"); } debug!(host = %host, addresses = ?addrs, "dns lookup successful"); info!(host = %host, addr = %addrs[0], "dns lookup result"); Ok(Resolution::Dns { host: host.to_string(), selected: addrs[0], addresses: addrs, }) } pub async fn resolve_ts3(&self, host: &str, proto: &str) -> Result { self.resolve_srv("ts3", DEFAULT_TS3_PROTO, host, proto) .await } pub async fn resolve_tsdns(&self, host: &str, proto: &str) -> Result { self.resolve_srv("tsdns", DEFAULT_TSDNS_PROTO, host, proto) .await } pub async fn resolve_nick(&self, nick: &str, proto: &str) -> Result { let proto = if proto.is_empty() { DEFAULT_TS3_PROTO } else { proto }; if proto == DEFAULT_TS3_PROTO { debug!(protocol = %proto, "using default protocol for nick srv lookup"); } info!(nick = %nick, service = "nick", "nick lookup started"); let addr = self.lookup_nick(nick).await?; if addr.is_empty() { bail!("no address found for nick: {nick}"); } debug!(nick = %nick, address = %addr, "nick lookup successful"); let records = self .lookup_srv("ts3", proto, &addr) .await .context("ts3 srv lookup for nick failed")?; if records.is_empty() { bail!("no srv records found for nick: {nick}"); } debug!(nick = %nick, srv_records = ?records, "ts3 srv lookup for nick successful"); info!( nick = %nick, resolved = %records[0].resolved(), "nick lookup result" ); Ok(Resolution::Nick { nick: nick.to_string(), address: addr, protocol: proto.to_string(), selected: records[0].clone(), records, }) } async fn resolve_srv( &self, service: &str, default_proto: &str, host: &str, proto: &str, ) -> Result { let proto = if proto.is_empty() { default_proto } else { proto }; if proto == default_proto { debug!(protocol = %proto, "using default protocol for srv lookup"); } info!( host = %host, service = %service, protocol = %proto, "{} srv lookup started", service ); let records = self.lookup_srv(service, proto, host).await?; if records.is_empty() { if service == "ts3" { bail!("no SRV records found for host: {host}"); } bail!("no srv records found for host: {host}"); } debug!( host = %host, srv_records = ?records, "{} srv lookup successful", service ); info!( host = %host, resolved = %records[0].resolved(), "{} srv lookup result", service ); Ok(Resolution::Srv { service: service.to_string(), host: host.to_string(), protocol: proto.to_string(), selected: records[0].clone(), records, }) } async fn lookup_srv(&self, service: &str, proto: &str, host: &str) -> Result> { if host.is_empty() { bail!("host cannot be empty"); } if service == "ts3" && proto == "tcp" { bail!("srv lookup for ts3 over tcp is not supported"); } let query = format!("_{service}._{proto}.{host}"); let resolver = self.resolver.as_ref().context("srv resolver unavailable")?; let result = resolver .srv_lookup(query) .await .context("srv lookup failed")?; let mut records = result .answers() .iter() .filter_map(|record| { record.try_borrow::().map(|record| { let data = record.data(); SrvRecord { priority: data.priority, weight: data.weight, port: data.port, target: data.target.to_string(), } }) }) .collect::>(); sort_srv_records(&mut records); Ok(records) } async fn lookup_ips(&self, host: &str, port: u16) -> Result> { let lookup = tokio::net::lookup_host((host, port)) .await .context("dns lookup failed")?; Ok(collect_ips(lookup)) } async fn lookup_nick(&self, nick: &str) -> Result { if nick.is_empty() { bail!("host cannot be empty"); } let value = lookup_named_value(&self.http, nick, false, "nick") .await? .ok_or_else(|| anyhow::anyhow!("no data returned for nick: {nick}"))?; if value.trim().is_empty() { bail!("no data returned for nick: {nick}"); } Ok(value) } } fn srv_resolver_builder() -> Result> { TokioResolver::builder_tokio().context("failed to initialize DNS resolver") } pub fn normalize_args(mut args: Args) -> Args { args.host = args.host.trim().to_lowercase(); args.service = args.service.trim().to_lowercase(); args.protocol = args.protocol.trim().to_lowercase(); args } pub fn validate_args(args: &Args) -> Result<()> { if args.host.is_empty() { bail!("host cannot be empty"); } match args.service.as_str() { "dns" | "ts3" | "tsdns" | "nick" => {} "" => bail!("service cannot be empty"), other => bail!("unsupported service: {other}"), } match args.protocol.as_str() { "" | "udp" | "tcp" => {} other => bail!("unsupported protocol: {other}"), } Ok(()) } pub async fn run(args: Args) -> Result<()> { let client = ChanoraResolver::new()?; let normalized = normalize_args(args); validate_args(&normalized)?; debug!(build = %build_info(), "app start"); debug!( host = %normalized.host, service = %normalized.service, protocol = %normalized.protocol, "normalized flags" ); client.resolve_normalized(&normalized).await.map(|_| ()) } fn collect_ips(lookup: impl Iterator) -> Vec { let mut addrs = lookup.map(|addr| addr.ip()).collect::>(); addrs.sort_by_key(|addr| match addr { IpAddr::V4(_) => 0u8, IpAddr::V6(_) => 1u8, }); addrs.dedup(); addrs } fn sort_srv_records(records: &mut [SrvRecord]) { records.sort_by(|left, right| { left.priority .cmp(&right.priority) .then_with(|| right.weight.cmp(&left.weight)) .then_with(|| left.target.cmp(&right.target)) .then_with(|| left.port.cmp(&right.port)) }); } fn format_host_port(host: &str, port: u16) -> String { if host.contains(':') && !host.starts_with('[') { format!("[{host}]:{port}") } else { format!("{host}:{port}") } } fn should_return_direct_dns_before_discovery(port: Option) -> bool { port.is_some() } fn tsdns_discovery_budget(has_dns_fallback: bool) -> Option { has_dns_fallback.then_some(TSDNS_FALLBACK_DISCOVERY_BUDGET) } fn ts3_srv_discovery_budget(has_dns_fallback: bool) -> Option { has_dns_fallback.then_some(TS3_SRV_FALLBACK_DISCOVERY_BUDGET) } fn direct_client_resolution( input: &str, from_server_name: bool, fallback_port: u16, resolution: Resolution, ) -> ClientResolution { let selected = match &resolution { Resolution::Dns { selected, .. } => *selected, _ => unreachable!("direct fallback must be a dns resolution"), }; client_resolution_with_address( input, if from_server_name { ClientResolutionMethod::Nick } else { ClientResolutionMethod::Dns }, format_host_port(&selected.to_string(), fallback_port), resolution, ) } fn client_resolution_with_address( input: &str, method: ClientResolutionMethod, address: String, resolution: Resolution, ) -> ClientResolution { client_resolution_with_optional_resolution(input, method, address, Some(resolution)) } fn client_resolution_with_optional_resolution( input: &str, method: ClientResolutionMethod, address: String, resolution: Option, ) -> ClientResolution { info!( input = %input, method = ?method, addr = %address, "client request resolution result" ); ClientResolution { input: input.to_string(), address, method, resolution, } } #[derive(Debug, Clone, PartialEq, Eq)] struct ClientRequest { host: String, port: Option, } #[derive(Debug, Clone, PartialEq, Eq)] struct ServerNameResolution { host: String, port: PortResolution, } #[derive(Debug, Clone, PartialEq, Eq)] enum PortResolution { Explicit(u16), KeepInput, } impl PortResolution { fn resolve(&self, fallback: u16) -> u16 { match self { PortResolution::Explicit(port) => *port, PortResolution::KeepInput => fallback, } } } fn parse_resolved_server_name(value: &str) -> Result { let value = value.trim().trim_matches('/'); if value.is_empty() { bail!("resolved server name cannot be empty"); } let value = strip_ts3server_scheme(value); let (authority, query_port) = split_authority_and_port(value)?; let authority = authority.trim_matches('/'); if authority.is_empty() { bail!("resolved server name cannot be empty"); } let authority = authority.trim().to_lowercase(); let (host, explicit_port) = match parse_host_port_components(&authority, true)? { Some((host, port)) => (host, port), None => ( authority, query_port .map(PortResolution::Explicit) .unwrap_or(PortResolution::KeepInput), ), }; Ok(ServerNameResolution { host, port: match explicit_port { PortResolution::KeepInput => query_port .map(PortResolution::Explicit) .unwrap_or(PortResolution::KeepInput), explicit => explicit, }, }) } fn should_try_server_name(host: &str) -> bool { !host.contains('.') && host.parse::().is_err() } fn tsdns_candidate_hosts(host: &str) -> Vec { let labels = host .trim_end_matches('.') .split('.') .filter(|label| !label.is_empty()) .collect::>(); if labels.is_empty() { return Vec::new(); } let mut candidates = Vec::new(); if labels.len() >= 2 { candidates.push(labels[labels.len() - 2..].join(".")); } let full = labels.join("."); if !candidates.iter().any(|candidate| candidate == &full) { candidates.push(full); } candidates } async fn lookup_named_value( http: &Client, name: &str, allow_not_found: bool, label: &str, ) -> Result> { let response = http .get(NICK_RESOLVE_URL) .query(&[("name", name)]) .send() .await .with_context(|| format!("{label} lookup failed"))?; read_lookup_body(response, allow_not_found, label).await } async fn read_lookup_body( response: reqwest::Response, allow_not_found: bool, label: &str, ) -> Result> { if allow_not_found && response.status() == reqwest::StatusCode::NOT_FOUND { return Ok(None); } if !response.status().is_success() { bail!("unexpected status code: {}", response.status()); } let bytes = response .bytes() .await .with_context(|| format!("{label} lookup failed"))?; if bytes.is_empty() { return Ok(None); } let value = String::from_utf8(bytes.to_vec()) .with_context(|| format!("{label} lookup returned invalid utf-8"))?; Ok(Some(value)) } fn normalize_client_input(input: &str) -> Result { let value = input.trim(); if value.is_empty() { bail!("host cannot be empty"); } let value = strip_ts3server_scheme(value); let (authority, query_port) = split_authority_and_port(value)?; let authority = authority.trim_matches('/'); if authority.is_empty() { bail!("host cannot be empty"); } let authority = authority.trim().to_lowercase(); let (host, explicit_port) = parse_host_port_components(&authority, false)? .unwrap_or((authority, PortResolution::KeepInput)); Ok(ClientRequest { host, port: match explicit_port { PortResolution::Explicit(port) => Some(port), PortResolution::KeepInput => query_port, }, }) } fn strip_ts3server_scheme(value: &str) -> &str { value .get(..12) .filter(|prefix| prefix.eq_ignore_ascii_case("ts3server://")) .map(|_| &value[12..]) .unwrap_or(value) } fn split_authority_and_port(value: &str) -> Result<(&str, Option)> { let Some((authority, query)) = value.split_once('?') else { return Ok((value, None)); }; Ok((authority, query_port(query)?)) } fn query_port(query: &str) -> Result> { for pair in query.split('&') { let Some((key, value)) = pair.split_once('=') else { continue; }; if key.eq_ignore_ascii_case("port") { return Ok(Some(parse_port(value)?)); } } Ok(None) } fn parse_host_port_components( input: &str, allow_keep_input: bool, ) -> Result> { if let Some(rest) = input.strip_prefix('[') { let Some(end) = rest.find(']') else { bail!("invalid bracketed ipv6 address"); }; let host = &rest[..end]; let suffix = &rest[end + 1..]; if suffix.is_empty() { return Ok(Some((host.to_string(), PortResolution::KeepInput))); } let Some(port) = suffix.strip_prefix(':') else { bail!("invalid bracketed ipv6 address"); }; return Ok(Some(( host.to_string(), parse_resolved_port(port, allow_keep_input)?, ))); } let Some((host, port)) = input.rsplit_once(':') else { return Ok(None); }; if host.contains(':') { return Ok(None); } Ok(Some(( host.to_string(), parse_resolved_port(port, allow_keep_input)?, ))) } fn is_bare_numeric_name(host: &str) -> bool { !host.is_empty() && host.chars().all(|value| value.is_ascii_digit()) } fn parse_port(value: &str) -> Result { if value.is_empty() { bail!("port cannot be empty"); } value .parse::() .with_context(|| format!("invalid port: {value}")) } fn parse_resolved_port(value: &str, allow_keep_input: bool) -> Result { if value.eq_ignore_ascii_case("$PORT") { if !allow_keep_input { bail!("port placeholder is not allowed in client input"); } return Ok(PortResolution::KeepInput); } Ok(PortResolution::Explicit(parse_port(value)?)) } #[cfg(test)] mod tests { use super::*; #[test] fn normalize_args_trims_and_lowercases() { let args = normalize_args(Args { host: " Example.COM ".into(), service: " DNS ".into(), protocol: " TCP ".into(), }); assert_eq!( args, Args { host: "example.com".into(), service: "dns".into(), protocol: "tcp".into(), } ); } #[test] fn validate_args_rejects_bad_service_and_proto() { let err = validate_args(&Args { host: "example.com".into(), service: "bad".into(), protocol: "".into(), }) .unwrap_err(); assert!(err.to_string().contains("unsupported service")); let err = validate_args(&Args { host: "example.com".into(), service: "dns".into(), protocol: "bad".into(), }) .unwrap_err(); assert!(err.to_string().contains("unsupported protocol")); } #[test] fn srv_records_sort_by_preferred_order() { let mut records = vec![ SrvRecord { priority: 20, weight: 10, port: 9987, target: "late.example.com.".into(), }, SrvRecord { priority: 10, weight: 5, port: 9987, target: "fallback.example.com.".into(), }, SrvRecord { priority: 10, weight: 20, port: 9987, target: "voice.example.com.".into(), }, ]; sort_srv_records(&mut records); assert_eq!(records[0].target, "voice.example.com."); assert_eq!(records[1].target, "fallback.example.com."); assert_eq!(records[2].target, "late.example.com."); } #[test] fn resolution_exposes_connection_address() { let dns = Resolution::Dns { host: "localhost".into(), addresses: vec!["::1".parse().unwrap()], selected: "::1".parse().unwrap(), }; assert_eq!(dns.connection_address(), "[::1]:9987"); let srv = Resolution::Srv { service: "ts3".into(), host: "example.com".into(), protocol: "udp".into(), records: vec![SrvRecord { priority: 0, weight: 0, port: 9988, target: "voice.example.com.".into(), }], selected: SrvRecord { priority: 0, weight: 0, port: 9988, target: "voice.example.com.".into(), }, }; assert_eq!(srv.connection_address(), "voice.example.com:9988"); } #[tokio::test] async fn tsdns_endpoint_preserves_srv_method_metadata() { let resolver = ChanoraResolver { resolver: None, http: Client::new(), }; let endpoint = ServerNameResolution { host: "127.0.0.1".into(), port: PortResolution::Explicit(9988), }; let srv_resolution = Resolution::Srv { service: "tsdns".into(), host: "example.com".into(), protocol: "tcp".into(), records: vec![SrvRecord { priority: 0, weight: 0, port: TSDNS_PORT, target: "tsdns.example.com.".into(), }], selected: SrvRecord { priority: 0, weight: 0, port: TSDNS_PORT, target: "tsdns.example.com.".into(), }, }; let via_srv = resolver .client_resolved_tsdns_endpoint( "example.com", endpoint.clone(), DEFAULT_TEAMSPEAK_PORT, Some(srv_resolution), ) .await .unwrap(); assert_eq!(via_srv.method, ClientResolutionMethod::TsdnsSrv); let via_tcp = resolver .client_resolved_tsdns_endpoint("example.com", endpoint, DEFAULT_TEAMSPEAK_PORT, None) .await .unwrap(); assert_eq!(via_tcp.method, ClientResolutionMethod::TsdnsTcp); } #[test] fn normalize_client_input_handles_chanora_request_shapes() { assert_eq!( normalize_client_input(" Voice.Example.com ").unwrap(), ClientRequest { host: "voice.example.com".into(), port: None, } ); assert_eq!( normalize_client_input("voice.example.com:9988").unwrap(), ClientRequest { host: "voice.example.com".into(), port: Some(9988), } ); assert_eq!( normalize_client_input("ts3server://voice.example.com?port=9989").unwrap(), ClientRequest { host: "voice.example.com".into(), port: Some(9989), } ); assert_eq!( normalize_client_input("[::1]:9987").unwrap(), ClientRequest { host: "::1".into(), port: Some(9987), } ); assert_eq!( normalize_client_input("[::1]").unwrap(), ClientRequest { host: "::1".into(), port: None, } ); } #[test] fn bare_numeric_names_are_not_dns_hosts() { assert!(is_bare_numeric_name("6666")); assert!(!is_bare_numeric_name("127.0.0.1")); assert!(!is_bare_numeric_name("voice6666")); } #[test] fn parse_resolved_server_name_accepts_host_and_port_shapes() { assert_eq!( parse_resolved_server_name("teamspeak.app").unwrap(), ServerNameResolution { host: "teamspeak.app".into(), port: PortResolution::KeepInput, } ); assert_eq!( parse_resolved_server_name("185.250.249.77:10075").unwrap(), ServerNameResolution { host: "185.250.249.77".into(), port: PortResolution::Explicit(10075), } ); assert_eq!( parse_resolved_server_name("ts3server://Voice.TeamSpeak.com?port=9989").unwrap(), ServerNameResolution { host: "voice.teamspeak.com".into(), port: PortResolution::Explicit(9989), } ); assert_eq!( parse_resolved_server_name("voice.teamspeak.com:$PORT").unwrap(), ServerNameResolution { host: "voice.teamspeak.com".into(), port: PortResolution::KeepInput, } ); } #[test] fn server_name_lookup_is_only_for_dotless_hosts() { assert!(should_try_server_name("6666")); assert!(should_try_server_name("wwb")); assert!(!should_try_server_name("voice.teamspeak.com")); assert!(!should_try_server_name("127.0.0.1")); } #[test] fn tsdns_candidates_include_parent_then_full_host() { assert_eq!( tsdns_candidate_hosts("voice.teamspeak.com"), vec!["teamspeak.com", "voice.teamspeak.com"] ); assert_eq!( tsdns_candidate_hosts("teamspeak.app"), vec!["teamspeak.app"] ); } #[test] fn explicit_ports_use_direct_dns_fast_path() { assert!(should_return_direct_dns_before_discovery(Some(9987))); assert!(!should_return_direct_dns_before_discovery(None)); } #[tokio::test] async fn explicit_port_client_request_uses_dns_result() { let resolver = ChanoraResolver { resolver: None, http: Client::new(), }; let resolved = resolver .resolve_client_request("localhost:10075") .await .unwrap(); assert_eq!(resolved.method, ClientResolutionMethod::Dns); assert!(resolved.address.ends_with(":10075")); } #[test] fn tsdns_discovery_is_bounded_when_dns_fallback_exists() { assert_eq!( tsdns_discovery_budget(true), Some(TSDNS_FALLBACK_DISCOVERY_BUDGET) ); assert_eq!(tsdns_discovery_budget(false), None); } #[test] fn ts3_srv_discovery_is_bounded_when_dns_fallback_exists() { assert_eq!( ts3_srv_discovery_budget(true), Some(TS3_SRV_FALLBACK_DISCOVERY_BUDGET) ); assert_eq!(ts3_srv_discovery_budget(false), None); } }