Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
4 changes: 3 additions & 1 deletion crates/openwire-core/src/body.rs
Original file line number Diff line number Diff line change
Expand Up @@ -243,7 +243,9 @@ impl ResponseBody {

pub async fn text(self) -> Result<String, WireError> {
let bytes = self.bytes().await?;
String::from_utf8(bytes.to_vec())
// Prefer reclaiming the buffer when this is the unique owner of the
// Bytes allocation (common after a full collect).
String::from_utf8(bytes.into())
.map_err(|error| WireError::body("response body is not valid UTF-8", error))
}

Expand Down
199 changes: 199 additions & 0 deletions crates/openwire-tokio/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -294,6 +294,145 @@ impl DnsResolver for SystemDnsResolver {
}
}

/// In-process DNS cache layered over another [`DnsResolver`].
///
/// Positive results are retained for `positive_ttl`. Empty / failed lookups are
/// retained for `negative_ttl` so transient NXDOMAIN storms do not hammer the
/// system resolver. Cache keys are `(host, port)`.
#[derive(Clone, Debug)]
pub struct CachingDnsResolver<R = SystemDnsResolver> {
inner: R,
positive_ttl: Duration,
negative_ttl: Duration,
cache: std::sync::Arc<std::sync::Mutex<std::collections::HashMap<DnsCacheKey, DnsCacheEntry>>>,
}

#[derive(Clone, Debug, PartialEq, Eq, Hash)]
struct DnsCacheKey {
host: String,
port: u16,
}

#[derive(Clone, Debug)]
struct DnsCacheEntry {
expires_at: Instant,
outcome: DnsCacheOutcome,
}

#[derive(Clone, Debug)]
enum DnsCacheOutcome {
Ok(Vec<SocketAddr>),
Err(String),
}

impl CachingDnsResolver<SystemDnsResolver> {
/// Builds a cache over the system resolver with 30s positive / 5s negative TTL.
pub fn system() -> Self {
Self::new(
SystemDnsResolver,
Duration::from_secs(30),
Duration::from_secs(5),
)
}
}

impl<R> CachingDnsResolver<R> {
pub fn new(inner: R, positive_ttl: Duration, negative_ttl: Duration) -> Self {
Self {
inner,
positive_ttl,
negative_ttl,
cache: std::sync::Arc::new(std::sync::Mutex::new(std::collections::HashMap::new())),
}
}

pub fn positive_ttl(&self) -> Duration {
self.positive_ttl
}

pub fn negative_ttl(&self) -> Duration {
self.negative_ttl
}
}

impl<R> DnsResolver for CachingDnsResolver<R>
where
R: DnsResolver + Clone,
{
fn resolve(
&self,
ctx: CallContext,
host: String,
port: u16,
) -> BoxFuture<Result<Vec<SocketAddr>, WireError>> {
let inner = self.inner.clone();
let cache = self.cache.clone();
let positive_ttl = self.positive_ttl;
let negative_ttl = self.negative_ttl;
Box::pin(async move {
let key = DnsCacheKey {
host: host.clone(),
port,
};
let now = Instant::now();
if let Ok(guard) = cache.lock() {
if let Some(entry) = guard.get(&key) {
if entry.expires_at > now {
match &entry.outcome {
DnsCacheOutcome::Ok(addrs) => {
ctx.listener().dns_start(&ctx, &host, port);
ctx.listener().dns_end(&ctx, &host, addrs);
return Ok(addrs.clone());
}
DnsCacheOutcome::Err(message) => {
ctx.listener().dns_start(&ctx, &host, port);
let error = WireError::dns(
message.clone(),
io::Error::new(io::ErrorKind::NotFound, message.clone()),
);
ctx.listener().dns_failed(&ctx, &host, &error);
return Err(error);
}
}
}
}
}

match inner.resolve(ctx.clone(), host.clone(), port).await {
Ok(addrs) => {
if let Ok(mut guard) = cache.lock() {
guard.insert(
key,
DnsCacheEntry {
expires_at: Instant::now() + positive_ttl,
outcome: DnsCacheOutcome::Ok(addrs.clone()),
},
);
// Opportunistic sweep of a few expired entries.
if guard.len() > 256 {
let now = Instant::now();
guard.retain(|_, entry| entry.expires_at > now);
}
}
Ok(addrs)
}
Err(error) => {
if let Ok(mut guard) = cache.lock() {
guard.insert(
key,
DnsCacheEntry {
expires_at: Instant::now() + negative_ttl,
outcome: DnsCacheOutcome::Err(error.message().to_owned()),
},
);
}
Err(error)
}
}
})
}
}

#[derive(Clone, Debug, Default)]
pub struct TokioTcpConnector;

Expand Down Expand Up @@ -434,4 +573,64 @@ mod tests {
timer.reset(&mut sleep, timer.now() + Duration::from_millis(1));
sleep.await;
}

#[tokio::test]
async fn caching_dns_resolver_reuses_positive_results() {
use std::net::{IpAddr, Ipv4Addr, SocketAddr};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;

use openwire_core::{
BoxFuture, CallContext, DnsResolver, NoopEventListener, SharedEventListener, WireError,
};

use super::CachingDnsResolver;

#[derive(Clone)]
struct CountingResolver {
hits: Arc<AtomicUsize>,
addr: SocketAddr,
}

impl DnsResolver for CountingResolver {
fn resolve(
&self,
ctx: CallContext,
host: String,
port: u16,
) -> BoxFuture<Result<Vec<SocketAddr>, WireError>> {
let hits = self.hits.clone();
let mut addr = self.addr;
Box::pin(async move {
hits.fetch_add(1, Ordering::Relaxed);
ctx.listener().dns_start(&ctx, &host, port);
addr.set_port(port);
let addrs = vec![addr];
ctx.listener().dns_end(&ctx, &host, &addrs);
Ok(addrs)
})
}
}

let hits = Arc::new(AtomicUsize::new(0));
let inner = CountingResolver {
hits: hits.clone(),
addr: SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 0, 2, 10)), 0),
};
let resolver =
CachingDnsResolver::new(inner, Duration::from_secs(60), Duration::from_secs(1));
let listener: SharedEventListener = Arc::new(NoopEventListener);
let ctx = CallContext::new(listener, None);

let first = resolver
.resolve(ctx.clone(), "example.com".into(), 80)
.await
.expect("first");
let second = resolver
.resolve(ctx, "example.com".into(), 80)
.await
.expect("second");
assert_eq!(first, second);
assert_eq!(hits.load(Ordering::Relaxed), 1);
}
}
75 changes: 29 additions & 46 deletions crates/openwire/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,7 @@ use openwire_core::{
RequestBody, ResponseBody, RetryPolicy, SharedEventListenerFactory, SharedInterceptor,
SharedTimer, TcpConnector, TlsConnector, WireError, WireExecutor,
};
use openwire_tokio::{SystemDnsResolver, TokioExecutor, TokioTcpConnector, TokioTimer};
use pin_project_lite::pin_project;
use openwire_tokio::{CachingDnsResolver, TokioExecutor, TokioTcpConnector, TokioTimer};
use tower::layer::Layer;
use tower::util::BoxCloneSyncService;
use tower::Service;
Expand Down Expand Up @@ -572,7 +571,7 @@ impl Default for ClientBuilder {
retry: RetryPolicyConfig::default(),
redirect: RedirectPolicyConfig::default(),
},
dns_resolver: Arc::new(SystemDnsResolver),
dns_resolver: Arc::new(CachingDnsResolver::system()),
tcp_connector: Arc::new(TokioTcpConnector),
tls_connector: None,
route_planner: Arc::new(DefaultRoutePlanner::default()),
Expand Down Expand Up @@ -987,32 +986,36 @@ impl QueuedCall {
}
}

/// Held on the response so admission capacity stays reserved for the body
/// lifetime without an extra `BoxBody` wrapper layer.
// RequestAdmissionPermit is not Clone; store behind Arc so it can live in
// response extensions until call lifecycle takes ownership.
#[derive(Clone)]
pub(crate) struct HeldRequestAdmission(pub(crate) std::sync::Arc<RequestAdmissionPermit>);

pub(crate) fn attach_request_admission(
response: Response<ResponseBody>,
mut response: Response<ResponseBody>,
permit: RequestAdmissionPermit,
) -> Response<ResponseBody> {
let (parts, body) = response.into_parts();
Response::from_parts(
parts,
ResponseBody::new(
RequestAdmissionBody {
inner: body,
_permit: Some(permit),
}
.boxed(),
),
)
response
.extensions_mut()
.insert(HeldRequestAdmission(std::sync::Arc::new(permit)));
response
}

fn attach_call_lifecycle(
response: Response<ResponseBody>,
mut response: Response<ResponseBody>,
ctx: CallContext,
state: Arc<CallState>,
) -> Response<ResponseBody> {
let admission = response
.extensions_mut()
.remove::<HeldRequestAdmission>()
.map(|held| held.0);
let (parts, body) = response.into_parts();
Response::from_parts(
parts,
ResponseBody::new(CallLifecycleBody::new(body, ctx, state).boxed()),
ResponseBody::new(CallLifecycleBody::new(body, ctx, state, admission).boxed()),
)
}

Expand Down Expand Up @@ -1338,15 +1341,23 @@ struct CallLifecycleBody {
inner: Option<ResponseBody>,
ctx: CallContext,
state: Arc<CallState>,
/// Keeps request admission capacity reserved until the body completes.
_admission: Option<std::sync::Arc<RequestAdmissionPermit>>,
finished: bool,
}

impl CallLifecycleBody {
fn new(inner: ResponseBody, ctx: CallContext, state: Arc<CallState>) -> Self {
fn new(
inner: ResponseBody,
ctx: CallContext,
state: Arc<CallState>,
admission: Option<std::sync::Arc<RequestAdmissionPermit>>,
) -> Self {
Self {
inner: Some(inner),
ctx,
state,
_admission: admission,
finished: false,
}
}
Expand Down Expand Up @@ -1440,34 +1451,6 @@ impl Body for CallLifecycleBody {
}
}

pin_project! {
struct RequestAdmissionBody {
#[pin]
inner: ResponseBody,
_permit: Option<RequestAdmissionPermit>,
}
}

impl Body for RequestAdmissionBody {
type Data = Bytes;
type Error = WireError;

fn poll_frame(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
self.project().inner.poll_frame(cx)
}

fn is_end_stream(&self) -> bool {
self.inner.is_end_stream()
}

fn size_hint(&self) -> SizeHint {
self.inner.size_hint()
}
}

#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
Expand Down
Loading
Loading