Skip to content
Merged
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 51 additions & 27 deletions server/svix-server/src/core/webhook_http_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ use ipnet::IpNet;
use openssl::ssl::{SslConnector, SslConnectorBuilder, SslMethod, SslVerifyMode};
use serde::Serialize;
use thiserror::Error;
use tokio::{net::TcpStream, sync::Mutex};
use tokio::{net::TcpStream, sync::OnceCell};
use tower::Service;

use crate::{
Expand Down Expand Up @@ -561,21 +561,15 @@ type NonLocalHttpConnector = HttpConnector<NonLocalDnsResolver>;
/// Specific private subnets or domain names may be whitelisted.
#[derive(Clone, Debug)]
struct NonLocalDnsResolver {
state: Arc<Mutex<DnsState>>,
resolver: Arc<OnceCell<TokioResolver>>,
whitelist_nets: Arc<Vec<IpNet>>,
whitelist_names: Arc<Vec<String>>,
}

#[derive(Clone, Debug)]
enum DnsState {
Init,
Ready(Arc<TokioResolver>),
}

impl NonLocalDnsResolver {
pub fn new(whitelist_nets: Arc<Vec<IpNet>>, whitelist_names: Arc<Vec<String>>) -> Self {
NonLocalDnsResolver {
state: Arc::new(Mutex::new(DnsState::Init)),
resolver: Arc::new(OnceCell::new()),
whitelist_nets,
whitelist_names,
}
Expand All @@ -592,24 +586,15 @@ impl Service<Name> for NonLocalDnsResolver {
}

fn call(&mut self, name: Name) -> Self::Future {
let resolver = self.clone();
let this = self.clone();
let whitelist_nets = self.whitelist_nets.clone();
let whitelist_names = self.whitelist_names.clone();

Box::pin(async move {
let mut lock = resolver.state.lock().await;

let resolver = match &*lock {
DnsState::Init => {
let resolver = new_resolver().await?;
*lock = DnsState::Ready(resolver.clone());
resolver
}

DnsState::Ready(resolver) => resolver.clone(),
};

drop(lock);
let resolver = this
.resolver
.get_or_try_init(|| async { new_resolver().await })
Comment thread
svix-jplatte marked this conversation as resolved.
Outdated
.await?;

let whitelisted_name = whitelist_names
.iter()
Expand Down Expand Up @@ -653,10 +638,10 @@ impl Iterator for SocketAddrs {
}
}

async fn new_resolver() -> Result<Arc<TokioResolver>, NetError> {
async fn new_resolver() -> Result<TokioResolver, NetError> {
let mut builder = Resolver::builder_tokio()?;
builder.options_mut().ip_strategy = hickory_resolver::config::LookupIpStrategy::Ipv4thenIpv6;
Ok(Arc::new(builder.build()?))
builder.build()
}

fn is_allowed(addr: IpAddr) -> bool {
Expand Down Expand Up @@ -719,15 +704,18 @@ mod tests {
net::{IpAddr, TcpListener},
path::PathBuf,
str::FromStr,
sync::Arc,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
};

use axum::{Router, routing};
use axum_server::tls_openssl::{OpenSSLAcceptor, OpenSSLConfig};
use http::{HeaderValue, Method, Version, header::AUTHORIZATION};
use ipnet::IpNet;

use super::{RequestBuilder, WebhookClient, is_allowed};
use super::{NonLocalDnsResolver, RequestBuilder, WebhookClient, is_allowed, new_resolver};
use crate::core::types::CasePreservingHeaderMap;

#[test]
Expand Down Expand Up @@ -942,4 +930,40 @@ mod tests {
let whc_without_validation = WebhookClient::new(Some(whitelist), None, true, None);
assert!(whc_without_validation.execute(request).await.is_ok());
}

#[tokio::test]
async fn test_dns_resolver_initializes_once_under_concurrency() {
Comment thread
svix-jplatte marked this conversation as resolved.
Outdated
let resolver = NonLocalDnsResolver::new(Arc::new(vec![]), Arc::new(vec![]));
let init_count = Arc::new(AtomicUsize::new(0));

// Spawn 20 tasks all racing to initialize the resolver at the same time
let handles: Vec<_> = (0..20)
.map(|_| {
let cell = resolver.resolver.clone();
let count = init_count.clone();
tokio::spawn(async move {
cell.get_or_try_init(|| async {
count.fetch_add(1, Ordering::SeqCst);
new_resolver().await
})
.await
.map(|_| ())
})
})
.collect();

for handle in handles {
handle.await.unwrap().expect("resolver init failed");
}

assert_eq!(
init_count.load(Ordering::SeqCst),
1,
"resolver should initialize exactly once regardless of concurrent callers"
);
assert!(
resolver.resolver.get().is_some(),
"OnceCell should be populated after initialization"
);
}
}
Loading