diff --git a/src/service/resolver/actual.rs b/src/service/resolver/actual.rs index 5b4b5995aa..d86d289bba 100644 --- a/src/service/resolver/actual.rs +++ b/src/service/resolver/actual.rs @@ -15,7 +15,7 @@ use tuwunel_core::{Err, Result, debug, debug_info, debug_warn, err, error, trace use super::{ DestString, FedDest, cache::{CachedDest, CachedOverride, MAX_IPS}, - fed::{PortString, add_port_to_hostname, get_ip_with_port}, + fed::{add_port_to_hostname, get_ip_with_port, srv_url_dest}, }; #[derive(Clone, Debug)] @@ -29,6 +29,28 @@ impl ActualDest { pub(crate) fn to_string(&self) -> DestString { self.dest.https_string() } } +#[derive(Clone, Debug, Eq, PartialEq)] +pub(super) struct SrvOverridePlan { + pub(super) base_hostname: DestString, + pub(super) srv_hostname: DestString, + pub(super) srv_port: u16, + pub(super) url_dest: FedDest, +} + +/// SRV records are looked up for the base hostname, but the request URL host +/// must be the SRV target so TLS SNI matches the peer certificate. +#[inline] +pub(super) fn srv_override_plan(base_hostname: &str, overrider: &FedDest) -> SrvOverridePlan { + let srv_hostname = overrider.hostname(); + + SrvOverridePlan { + base_hostname: base_hostname.into(), + srv_hostname: srv_hostname.as_str().into(), + srv_port: overrider.port().unwrap_or(8448), + url_dest: srv_url_dest(overrider), + } +} + impl super::Service { #[tracing::instrument(skip_all, level = "debug", name = "resolve")] pub(crate) async fn get_actual_dest(&self, server_name: &ServerName) -> Result { @@ -196,26 +218,16 @@ impl super::Service { overrider: FedDest, ) -> Result { debug!("3.3: SRV lookup successful"); - let force_port = overrider.port(); + let srv = srv_override_plan(delegated, &overrider); self.conditional_query_and_cache_override( - delegated, - &overrider.hostname(), - force_port.unwrap_or(8448), + srv.base_hostname.as_str(), + srv.srv_hostname.as_str(), + srv.srv_port, cache, ) .await?; - if let Some(port) = force_port { - return Ok(FedDest::Named( - delegated.into(), - format!(":{port}") - .as_str() - .try_into() - .unwrap_or_else(|_| FedDest::default_port()), - )); - } - - Ok(add_port_to_hostname(delegated)) + Ok(srv.url_dest) } async fn actual_dest_3_4(&self, cache: bool, delegated: &str) -> Result { @@ -233,24 +245,16 @@ impl super::Service { overrider: FedDest, ) -> Result { debug!("4: No .well-known; SRV record found"); - let force_port = overrider.port(); + let srv = srv_override_plan(host, &overrider); self.conditional_query_and_cache_override( - host, - &overrider.hostname(), - force_port.unwrap_or(8448), + srv.base_hostname.as_str(), + srv.srv_hostname.as_str(), + srv.srv_port, cache, ) .await?; - if let Some(port) = force_port { - let port = format!(":{port}"); - return Ok(FedDest::Named( - host.into(), - PortString::from(port.as_str()).unwrap_or_else(|_| FedDest::default_port()), - )); - } - - Ok(add_port_to_hostname(host)) + Ok(srv.url_dest) } async fn actual_dest_5(&self, dest: &ServerName, cache: bool) -> Result { diff --git a/src/service/resolver/fed.rs b/src/service/resolver/fed.rs index 0cf9552e56..ea536bbedc 100644 --- a/src/service/resolver/fed.rs +++ b/src/service/resolver/fed.rs @@ -44,6 +44,21 @@ pub(crate) fn add_port_to_hostname(dest: &str) -> FedDest { ) } +/// URL hostname is the SRV target so TLS SNI matches the SRV peer's cert. +pub(crate) fn srv_url_dest(overrider: &FedDest) -> FedDest { + let target = overrider.hostname(); + match overrider.port() { + | Some(port) => FedDest::Named( + target.as_str().into(), + format!(":{port}") + .as_str() + .try_into() + .unwrap_or_else(|_| FedDest::default_port()), + ), + | None => add_port_to_hostname(target.as_str()), + } +} + impl FedDest { pub(crate) fn https_string(&self) -> DestString { match self { diff --git a/src/service/resolver/tests.rs b/src/service/resolver/tests.rs index 709365c7a7..c96d24a9aa 100644 --- a/src/service/resolver/tests.rs +++ b/src/service/resolver/tests.rs @@ -1,4 +1,7 @@ -use super::fed::{FedDest, add_port_to_hostname, get_ip_with_port}; +use super::{ + actual::srv_override_plan, + fed::{FedDest, add_port_to_hostname, get_ip_with_port}, +}; #[test] fn ips_get_default_ports() { @@ -39,3 +42,18 @@ fn hostnames_keep_custom_ports() { FedDest::Named("example.com".into(), ":1337".try_into().unwrap()) ); } + +#[test] +fn srv_override_plan_uses_srv_target_as_url_hostname() { + let srv = FedDest::Named("matrix.example.net".into(), ":443".try_into().unwrap()); + + let plan = srv_override_plan("example.com", &srv); + + assert_eq!(plan.base_hostname.as_str(), "example.com"); + assert_eq!(plan.srv_hostname.as_str(), "matrix.example.net"); + assert_eq!(plan.srv_port, 443); + assert_eq!( + plan.url_dest, + FedDest::Named("matrix.example.net".into(), ":443".try_into().unwrap()), + ); +}