From d20fdb24a2bb51c968b53fd9796e3fcd908a0869 Mon Sep 17 00:00:00 2001 From: fannnzhang Date: Sat, 11 Jul 2026 17:06:36 +0800 Subject: [PATCH 1/2] perf: shard the connection pool and cut hot-path allocations Reduce lock contention and cloning on the request hot path: - Shard the connection pool by address hash (32 shards) with a shared coalescing index so multi-host acquire/release no longer serializes on one global mutex. - Share a single Arc
across dual-stack / multi-endpoint route plans. - Store follow-up RequestSnapshot headers and extensions behind Arc so auth and retry rebuilds avoid repeated full-map clones when only metadata is used. - Reclaim ResponseBody::text() buffers via Bytes::into() when uniquely owned. --- crates/openwire-core/src/body.rs | 4 +- crates/openwire/src/connection/planning.rs | 45 ++- crates/openwire/src/connection/pool.rs | 318 ++++++++++++--------- crates/openwire/src/policy/follow_up.rs | 24 +- docs/ARCHITECTURE.md | 13 + 5 files changed, 250 insertions(+), 154 deletions(-) diff --git a/crates/openwire-core/src/body.rs b/crates/openwire-core/src/body.rs index a42bac2..c1c29e8 100644 --- a/crates/openwire-core/src/body.rs +++ b/crates/openwire-core/src/body.rs @@ -243,7 +243,9 @@ impl ResponseBody { pub async fn text(self) -> Result { 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)) } diff --git a/crates/openwire/src/connection/planning.rs b/crates/openwire/src/connection/planning.rs index a8d5b9c..a02f240 100644 --- a/crates/openwire/src/connection/planning.rs +++ b/crates/openwire/src/connection/planning.rs @@ -328,7 +328,7 @@ pub(crate) enum RouteKind { #[derive(Clone, Debug, PartialEq, Eq, Hash)] pub struct Route { - address: Address, + address: Arc
, family: RouteFamily, kind: RouteKind, target_dns: DnsResolution, @@ -337,6 +337,17 @@ pub struct Route { impl Route { pub fn direct(address: Address, target: SocketAddr) -> Self { + Self { + family: RouteFamily::from_socket_addr(target), + kind: RouteKind::Direct { target }, + target_dns: DnsResolution::Local(target), + proxy_dns: None, + address: Arc::new(address), + } + } + + /// Builds a direct route reusing a shared address allocation (dual-stack plans). + pub(crate) fn direct_shared(address: Arc
, target: SocketAddr) -> Self { Self { family: RouteFamily::from_socket_addr(target), kind: RouteKind::Direct { target }, @@ -347,11 +358,15 @@ impl Route { } pub fn http_forward(address: Address, proxy: SocketAddr) -> Self { + Self::http_forward_shared(Arc::new(address), proxy) + } + + pub(crate) fn http_forward_shared(address: Arc
, proxy: SocketAddr) -> Self { let credentials = address .proxy() .and_then(|proxy| proxy.endpoint().credentials()) .cloned(); - Self::proxy_route( + Self::proxy_route_shared( address, proxy, RouteKind::HttpForwardProxy { proxy, credentials }, @@ -359,11 +374,15 @@ impl Route { } pub fn connect_proxy(address: Address, proxy: SocketAddr) -> Self { + Self::connect_proxy_shared(Arc::new(address), proxy) + } + + pub(crate) fn connect_proxy_shared(address: Arc
, proxy: SocketAddr) -> Self { let credentials = address .proxy() .and_then(|proxy| proxy.endpoint().credentials()) .cloned(); - Self::proxy_route( + Self::proxy_route_shared( address, proxy, RouteKind::ConnectProxy { proxy, credentials }, @@ -371,11 +390,15 @@ impl Route { } pub fn socks_proxy(address: Address, proxy: SocketAddr) -> Self { + Self::socks_proxy_shared(Arc::new(address), proxy) + } + + pub(crate) fn socks_proxy_shared(address: Arc
, proxy: SocketAddr) -> Self { let credentials = address .proxy() .and_then(|proxy| proxy.endpoint().credentials()) .cloned(); - Self::proxy_route(address, proxy, RouteKind::SocksProxy { proxy, credentials }) + Self::proxy_route_shared(address, proxy, RouteKind::SocksProxy { proxy, credentials }) } pub(crate) fn from_observed(address: Address, remote_addr: Option) -> Self { @@ -396,7 +419,7 @@ impl Route { } } - fn proxy_route(address: Address, proxy: SocketAddr, kind: RouteKind) -> Self { + fn proxy_route_shared(address: Arc
, proxy: SocketAddr, kind: RouteKind) -> Self { let authority = address.authority(); Self { family: RouteFamily::from_socket_addr(proxy), @@ -661,10 +684,11 @@ impl DefaultRoutePlanner { resolved_addrs: impl IntoIterator, ) -> RoutePlan { let ordered = order_for_fast_fallback(resolved_addrs); + let address = Arc::new(address); RoutePlan::new( ordered .into_iter() - .map(|addr| Route::direct(address.clone(), addr)) + .map(|addr| Route::direct_shared(address.clone(), addr)) .collect(), self.fast_fallback_stagger, ) @@ -676,10 +700,11 @@ impl DefaultRoutePlanner { resolved_proxy_addrs: impl IntoIterator, ) -> RoutePlan { let ordered = order_for_fast_fallback(resolved_proxy_addrs); + let address = Arc::new(address); RoutePlan::new( ordered .into_iter() - .map(|addr| Route::http_forward(address.clone(), addr)) + .map(|addr| Route::http_forward_shared(address.clone(), addr)) .collect(), self.fast_fallback_stagger, ) @@ -691,10 +716,11 @@ impl DefaultRoutePlanner { resolved_proxy_addrs: impl IntoIterator, ) -> RoutePlan { let ordered = order_for_fast_fallback(resolved_proxy_addrs); + let address = Arc::new(address); RoutePlan::new( ordered .into_iter() - .map(|addr| Route::connect_proxy(address.clone(), addr)) + .map(|addr| Route::connect_proxy_shared(address.clone(), addr)) .collect(), self.fast_fallback_stagger, ) @@ -706,10 +732,11 @@ impl DefaultRoutePlanner { resolved_proxy_addrs: impl IntoIterator, ) -> RoutePlan { let ordered = order_for_fast_fallback(resolved_proxy_addrs); + let address = Arc::new(address); RoutePlan::new( ordered .into_iter() - .map(|addr| Route::socks_proxy(address.clone(), addr)) + .map(|addr| Route::socks_proxy_shared(address.clone(), addr)) .collect(), self.fast_fallback_stagger, ) diff --git a/crates/openwire/src/connection/pool.rs b/crates/openwire/src/connection/pool.rs index cd489ea..257dbb2 100644 --- a/crates/openwire/src/connection/pool.rs +++ b/crates/openwire/src/connection/pool.rs @@ -1,4 +1,5 @@ use std::collections::{HashMap, HashSet}; +use std::hash::{Hash, Hasher}; use std::net::SocketAddr; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -15,6 +16,10 @@ use crate::sync_util::lock_mutex; /// abort the owned hyper task and drop bindings. Must not re-enter the pool. pub(crate) type PoolEvictionHook = Arc; +/// Number of address-keyed pool shards. Keeps independent hosts off the same +/// mutex while preserving exact-address reuse semantics within a shard. +const POOL_SHARDS: usize = 32; + #[derive(Clone, Debug, PartialEq, Eq)] pub(crate) struct PoolSettings { pub(crate) idle_timeout: Option, @@ -56,22 +61,30 @@ pub(crate) struct PoolStats { pub(crate) struct ConnectionPool { settings: PoolSettings, - state: Mutex, + shards: Arc<[Mutex]>, + /// Coalescing index is shared: candidates may live on different address + /// shards. Guarded separately from per-address shards. + coalesced_by_target: Mutex>>, + /// Global id → address map for remove-by-id without scanning shards. + by_id: Mutex>, eviction_hook: Mutex>, } #[derive(Debug, Default)] struct PoolState { by_address: HashMap>, - by_id: HashMap, - coalesced_by_target: HashMap>, } impl ConnectionPool { pub(crate) fn new(settings: PoolSettings) -> Self { + let shards = (0..POOL_SHARDS) + .map(|_| Mutex::new(PoolState::default())) + .collect::>(); Self { settings, - state: Mutex::new(PoolState::default()), + shards: Arc::<[Mutex]>::from(shards), + coalesced_by_target: Mutex::new(HashMap::new()), + by_id: Mutex::new(HashMap::new()), eviction_hook: Mutex::new(None), } } @@ -88,29 +101,36 @@ impl ConnectionPool { &self.settings } + fn shard(&self, address: &Address) -> &Mutex { + &self.shards[address_shard(address)] + } + pub(crate) fn insert(&self, connection: RealConnection) { let address = connection.address().clone(); - let mut state = lock_mutex(&self.state); - state - .by_address - .entry(address.clone()) - .or_default() - .push(connection.clone()); - register_connection(&mut state, &connection); - let evicted = prune_address(&self.settings, &mut state, &address); - drop(state); - self.notify_evictions(evicted); + let mut evicted = Vec::new(); + { + let mut state = lock_mutex(self.shard(&address)); + state + .by_address + .entry(address.clone()) + .or_default() + .push(connection.clone()); + evicted.extend(prune_address(&self.settings, &mut state, &address)); + } + lock_mutex(&self.by_id).insert(connection.id(), address); + self.index_coalescing(&connection); + self.finish_removals(evicted); } pub(crate) fn acquire(&self, address: &Address) -> Option { - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, address); + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); let result = state .by_address .get_mut(address) .and_then(|connections| connections.iter().find(|conn| conn.try_acquire()).cloned()); drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); result } @@ -118,8 +138,8 @@ impl ConnectionPool { &self, address: &Address, ) -> (Option, bool) { - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, address); + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); let result = match state.by_address.get_mut(address) { Some(connections) => { let connection = connections.iter().find(|conn| conn.try_acquire()).cloned(); @@ -132,19 +152,19 @@ impl ConnectionPool { None => (None, false), }; drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); result } pub(crate) fn has_in_use_connection(&self, address: &Address) -> bool { - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, address); + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); let result = state .by_address .get(address) .is_some_and(|connections| has_in_use_connection_unpruned(connections, None)); drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); result } @@ -162,43 +182,48 @@ impl ConnectionPool { return None; } - let mut state = lock_mutex(&self.state); let mut addresses_to_prune = HashSet::new(); - for target in &direct_targets { - if let Some(bucket) = state.coalesced_by_target.get(target) { - addresses_to_prune - .extend(bucket.iter().map(|connection| connection.address().clone())); + { + let coalesced = lock_mutex(&self.coalesced_by_target); + for target in &direct_targets { + if let Some(bucket) = coalesced.get(target) { + addresses_to_prune + .extend(bucket.iter().map(|connection| connection.address().clone())); + } } } - let mut evicted = Vec::new(); - for candidate_address in addresses_to_prune { - evicted.extend(prune_address( + let mut removed = Vec::new(); + for candidate_address in &addresses_to_prune { + let mut state = lock_mutex(self.shard(candidate_address)); + removed.extend(prune_address( &self.settings, &mut state, - &candidate_address, + candidate_address, )); } + self.finish_removals(removed); let mut candidates = Vec::new(); let mut seen_ids = HashSet::new(); - for target in direct_targets { - prune_coalescing_bucket(&mut state, target); - let Some(bucket) = state.coalesced_by_target.get(&target) else { - continue; - }; - - for connection in bucket { - if seen_ids.insert(connection.id()) && can_coalesce(connection, address, route_plan) - { - candidates.push(connection.clone()); + { + let mut coalesced = lock_mutex(&self.coalesced_by_target); + for target in direct_targets { + prune_coalescing_bucket(&mut coalesced, target); + let Some(bucket) = coalesced.get(&target) else { + continue; + }; + + for connection in bucket { + if seen_ids.insert(connection.id()) + && can_coalesce(connection, address, route_plan) + { + candidates.push(connection.clone()); + } } } } - drop(state); - self.notify_evictions(evicted); - candidates .into_iter() .find(|connection| connection.try_acquire()) @@ -209,11 +234,11 @@ impl ConnectionPool { return false; } - let address = connection.address().clone(); - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, &address); + let address = connection.address(); + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); true } @@ -222,9 +247,9 @@ impl ConnectionPool { address: &Address, connection_id: ConnectionId, ) -> Option { - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, address); - let result = if state.by_id.get(&connection_id) != Some(address) { + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); + let result = if lock_mutex(&self.by_id).get(&connection_id) != Some(address) { None } else { state.by_address.get(address).and_then(|connections| { @@ -235,7 +260,7 @@ impl ConnectionPool { }) }; drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); result } @@ -244,9 +269,9 @@ impl ConnectionPool { address: &Address, connection_id: ConnectionId, ) -> Option { - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, address); - let result = if state.by_id.get(&connection_id) != Some(address) { + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); + let result = if lock_mutex(&self.by_id).get(&connection_id) != Some(address) { None } else { state.by_address.get_mut(address).and_then(|connections| { @@ -257,7 +282,7 @@ impl ConnectionPool { }) }; drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); result } @@ -277,8 +302,8 @@ impl ConnectionPool { &self, connection_id: ConnectionId, ) -> Option { - let mut state = lock_mutex(&self.state); - let address = state.by_id.remove(&connection_id)?; + let address = lock_mutex(&self.by_id).remove(&connection_id)?; + let mut state = lock_mutex(self.shard(&address)); let mut removed = None; let mut should_remove_key = false; @@ -299,45 +324,75 @@ impl ConnectionPool { } if let Some(ref connection) = removed { - remove_index_connection(&mut state.coalesced_by_target, connection); + remove_index_connection(&mut lock_mutex(&self.coalesced_by_target), connection); } removed } pub(crate) fn stats(&self, address: &Address) -> PoolStats { - let mut state = lock_mutex(&self.state); - let evicted = prune_address(&self.settings, &mut state, address); + let mut state = lock_mutex(self.shard(address)); + let removed = prune_address(&self.settings, &mut state, address); let stats = match state.by_address.get(address) { - Some(connections) => { - connections - .iter() - .fold(PoolStats::default(), |mut stats, connection| { - stats.total += 1; - match connection.snapshot().allocation { - ConnectionAllocationState::Idle => stats.idle += 1, - ConnectionAllocationState::InUse { .. } => stats.in_use += 1, - ConnectionAllocationState::Closed => {} - } - stats - }) - } + Some(connections) => connections.iter().fold( + PoolStats::default(), + |mut stats, connection| { + stats.total += 1; + match connection.snapshot().allocation { + ConnectionAllocationState::Idle => stats.idle += 1, + ConnectionAllocationState::InUse { .. } => stats.in_use += 1, + ConnectionAllocationState::Closed => {} + } + stats + }, + ), None => PoolStats::default(), }; drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); stats } pub(crate) fn prune_all(&self) { - let mut state = lock_mutex(&self.state); - let addresses = state.by_address.keys().cloned().collect::>(); - let mut evicted = Vec::new(); - for address in addresses { - evicted.extend(prune_address(&self.settings, &mut state, &address)); + let mut removed = Vec::new(); + for shard in self.shards.iter() { + let mut state = lock_mutex(shard); + let addresses = state.by_address.keys().cloned().collect::>(); + for address in addresses { + removed.extend(prune_address(&self.settings, &mut state, &address)); + } } - drop(state); - self.notify_evictions(evicted); + self.finish_removals(removed); + } + + fn index_coalescing(&self, connection: &RealConnection) { + let Some(target) = coalescing_index_target(connection) else { + return; + }; + lock_mutex(&self.coalesced_by_target) + .entry(target) + .or_default() + .push(connection.clone()); + } + + fn finish_removals(&self, removed: Vec) { + if removed.is_empty() { + return; + } + let ids = removed.iter().map(RealConnection::id).collect::>(); + { + let mut by_id = lock_mutex(&self.by_id); + for connection in &removed { + by_id.remove(&connection.id()); + } + } + { + let mut coalesced = lock_mutex(&self.coalesced_by_target); + for connection in &removed { + remove_index_connection(&mut coalesced, connection); + } + } + self.notify_evictions(ids); } fn notify_evictions(&self, connection_ids: Vec) { @@ -362,12 +417,12 @@ impl std::fmt::Debug for ConnectionPool { } } -/// Returns connection ids closed and removed from the pool. +/// Returns connections closed and removed from the address bucket. fn prune_address( settings: &PoolSettings, state: &mut PoolState, address: &Address, -) -> Vec { +) -> Vec { let (removed, empty) = { let Some(connections) = state.by_address.get_mut(address) else { return Vec::new(); @@ -377,17 +432,17 @@ fn prune_address( (removed, connections.is_empty()) }; - let mut ids = Vec::with_capacity(removed.len()); - for connection in &removed { - unregister_connection(state, connection); - ids.push(connection.id()); - } - if empty { state.by_address.remove(address); } - ids + removed +} + +fn address_shard(address: &Address) -> usize { + let mut hasher = std::collections::hash_map::DefaultHasher::new(); + address.hash(&mut hasher); + (hasher.finish() as usize) % POOL_SHARDS } fn prune_connections( @@ -600,28 +655,6 @@ fn direct_route_targets(route_plan: &RoutePlan) -> Vec { targets } -fn register_connection(state: &mut PoolState, connection: &RealConnection) { - state - .by_id - .insert(connection.id(), connection.address().clone()); - index_connection(&mut state.coalesced_by_target, connection); -} - -fn unregister_connection(state: &mut PoolState, connection: &RealConnection) { - state.by_id.remove(&connection.id()); - remove_index_connection(&mut state.coalesced_by_target, connection); -} - -fn index_connection( - index: &mut HashMap>, - connection: &RealConnection, -) { - let Some(target) = coalescing_index_target(connection) else { - return; - }; - index.entry(target).or_default().push(connection.clone()); -} - fn remove_index_connection( index: &mut HashMap>, connection: &RealConnection, @@ -640,8 +673,11 @@ fn remove_index_connection( } } -fn prune_coalescing_bucket(state: &mut PoolState, target: SocketAddr) { - let should_remove = if let Some(bucket) = state.coalesced_by_target.get_mut(&target) { +fn prune_coalescing_bucket( + coalesced: &mut HashMap>, + target: SocketAddr, +) { + let should_remove = if let Some(bucket) = coalesced.get_mut(&target) { bucket.retain(|connection| { coalescing_index_target(connection) == Some(target) && !connection.is_closed() }); @@ -650,7 +686,7 @@ fn prune_coalescing_bucket(state: &mut PoolState, target: SocketAddr) { false }; if should_remove { - state.coalesced_by_target.remove(&target); + coalesced.remove(&target); } } @@ -763,7 +799,7 @@ mod tests { let connection_id = connection.id(); pool.insert(connection.clone()); assert_eq!( - lock_mutex(&pool.state).by_id.get(&connection_id), + lock_mutex(&pool.by_id).get(&connection_id), Some(&address) ); @@ -805,7 +841,7 @@ mod tests { let removed = pool.remove(connection_id).expect("connection should exist"); assert_eq!(removed.id(), connection_id); - assert!(!lock_mutex(&pool.state).by_id.contains_key(&connection_id)); + assert!(!lock_mutex(&pool.by_id).contains_key(&connection_id)); assert_eq!(pool.stats(&address), PoolStats::default()); } @@ -1057,7 +1093,7 @@ mod tests { assert!(pool.acquire(&address).is_none()); assert!(pool.get_by_id(&address, connection_id).is_none()); assert!(pool.remove(connection_id).is_none()); - assert!(!lock_mutex(&pool.state).by_id.contains_key(&connection_id)); + assert!(!lock_mutex(&pool.by_id).contains_key(&connection_id)); assert_eq!(connection.snapshot().health, ConnectionHealth::Closed); } @@ -1109,13 +1145,9 @@ mod tests { let pool = ConnectionPool::new(PoolSettings::default()); pool.insert(connection); - assert!(lock_mutex(&pool.state) - .coalesced_by_target - .contains_key(&target)); + assert!(lock_mutex(&pool.coalesced_by_target).contains_key(&target)); assert!(pool.remove(connection_id).is_some()); - assert!(!lock_mutex(&pool.state) - .coalesced_by_target - .contains_key(&target)); + assert!(!lock_mutex(&pool.coalesced_by_target).contains_key(&target)); } #[test] @@ -1139,10 +1171,11 @@ mod tests { pool.prune_all(); - let state = lock_mutex(&pool.state); - assert!(!state.by_address.contains_key(&address)); - assert!(!state.by_id.contains_key(&connection_id)); - assert!(!state.coalesced_by_target.contains_key(&target)); + assert!(!lock_mutex(pool.shard(&address)) + .by_address + .contains_key(&address)); + assert!(!lock_mutex(&pool.by_id).contains_key(&connection_id)); + assert!(!lock_mutex(&pool.coalesced_by_target).contains_key(&target)); } #[test] @@ -1174,9 +1207,11 @@ mod tests { in_use: 0, } ); - let state = lock_mutex(&pool.state); - assert!(!state.by_id.contains_key(&stale_id)); - assert_eq!(state.by_id.get(&live_id), Some(&live_address)); + assert!(!lock_mutex(&pool.by_id).contains_key(&stale_id)); + assert_eq!( + lock_mutex(&pool.by_id).get(&live_id), + Some(&live_address) + ); } #[test] @@ -1189,7 +1224,7 @@ mod tests { let _ = panic::catch_unwind(AssertUnwindSafe(|| { let _guard = pool - .state + .shard(&address) .lock() .expect("poison connection pool lock for test"); panic!("poison connection pool"); @@ -1319,4 +1354,19 @@ mod tests { assert!(!pool.has_in_use_connection(&address)); } + #[test] + fn different_hosts_map_to_independent_shards() { + let a = address_for_host("a.test", None, ProtocolPolicy::Http1OrHttp2); + let b = address_for_host("b.test", None, ProtocolPolicy::Http1OrHttp2); + assert_eq!(super::address_shard(&a), super::address_shard(&a)); + assert_eq!(super::address_shard(&b), super::address_shard(&b)); + let pool = ConnectionPool::new(PoolSettings::default()); + pool.insert(make_connection(a.clone(), 1)); + pool.insert(make_connection(b.clone(), 2)); + assert!(pool.acquire(&a).is_some()); + assert!(pool.acquire(&b).is_some()); + assert!(pool.acquire(&a).is_none()); + assert!(pool.acquire(&b).is_none()); + } + } diff --git a/crates/openwire/src/policy/follow_up.rs b/crates/openwire/src/policy/follow_up.rs index 15a2741..8c4b373 100644 --- a/crates/openwire/src/policy/follow_up.rs +++ b/crates/openwire/src/policy/follow_up.rs @@ -1,3 +1,4 @@ +use std::sync::Arc; use std::task::{Context, Poll}; use std::time::SystemTime; @@ -384,8 +385,8 @@ async fn authenticate_response( snapshot.method.clone(), snapshot.uri.clone(), snapshot.version, - snapshot.headers.clone(), - snapshot.extensions.clone(), + (*snapshot.headers).clone(), + (*snapshot.extensions).clone(), snapshot.body.as_ref().and_then(RequestBody::try_clone), ), AuthResponseState::new(response.status(), response.headers().clone()), @@ -511,8 +512,10 @@ struct RequestSnapshot { method: Method, uri: Uri, version: Version, - headers: HeaderMap, - extensions: http::Extensions, + /// Shared so auth / retry rebuilds do not re-clone the full map on each hop. + headers: Arc, + /// Shared request extensions captured at attempt start. + extensions: Arc, body: Option, } @@ -528,8 +531,8 @@ impl RequestSnapshot { method: request.method().clone(), uri: request.uri().clone(), version: request.version(), - headers: request.headers().clone(), - extensions: request.extensions().clone(), + headers: Arc::new(request.headers().clone()), + extensions: Arc::new(request.extensions().clone()), body: request.body().try_clone(), } } @@ -557,8 +560,8 @@ impl RequestSnapshot { .uri(self.uri.clone()) .version(self.version) .body(body)?; - *request.headers_mut() = self.headers.clone(); - *request.extensions_mut() = self.extensions.clone(); + *request.headers_mut() = (*self.headers).clone(); + *request.extensions_mut() = (*self.extensions).clone(); let sticky_proxy = request.extensions().get::().cloned(); reset_network_attempt_extensions(request.extensions_mut(), sticky_proxy); request.extensions_mut().insert(policy_trace); @@ -603,7 +606,7 @@ impl RequestSnapshot { self.method }; - let mut headers = self.headers; + let mut headers = Arc::try_unwrap(self.headers).unwrap_or_else(|arc| (*arc).clone()); headers.remove(HOST); if !same_origin { strip_sensitive_cross_origin_headers(&mut headers); @@ -622,7 +625,8 @@ impl RequestSnapshot { .version(self.version) .body(body)?; *request.headers_mut() = headers; - *request.extensions_mut() = self.extensions; + *request.extensions_mut() = + Arc::try_unwrap(self.extensions).unwrap_or_else(|arc| (*arc).clone()); reset_network_attempt_extensions(request.extensions_mut(), selected_proxy); request.extensions_mut().insert(policy_trace); Ok(Some(request)) diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 57b01a0..3687d8b 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -345,3 +345,16 @@ cargo test -p openwire --test live_network -- --ignored --test-threads=1 - Transparent decompression failures mark the call for connection discard so HTTP/2 connections are not returned to the pool as healthy after a body error. +## Performance notes + +- The connection pool is sharded by address hash (`32` shards) so concurrent + acquire/release for different hosts does not serialize on one global mutex. + HTTP/2 coalescing still uses a shared index keyed by direct route target. +- Dual-stack route plans share a single `Arc
` across candidate routes + instead of cloning the full address key per IP. +- Follow-up `RequestSnapshot` stores headers and extensions behind `Arc` so + auth challenge construction and retry rebuilds avoid re-cloning large maps + when only metadata is inspected. +- `ResponseBody::text()` reclaims the collected `Bytes` buffer via `into()` when + unique, avoiding an extra `to_vec()` copy on the common path. + From db625061b7f1ab9ee9d611907d3274fc6bfc0e63 Mon Sep 17 00:00:00 2001 From: fannnzhang Date: Sat, 11 Jul 2026 17:15:52 +0800 Subject: [PATCH 2/2] perf: lazy snapshots, DNS cache, and thinner response body path Build on the rebased pool sharding work: - Capture follow-up RequestSnapshot headers/extensions only when redirects, retries, or authenticators can rebuild the request (light mode for single-shot). - Default clients use CachingDnsResolver (30s positive / 5s negative TTL) over the system resolver to cut repeated lookups under connection churn. - Hold request-admission permits via response extensions into CallLifecycleBody, removing an extra BoxBody wrapper on returned responses. Also merges pool sharding with the #74 eviction-hook teardown path. --- crates/openwire-tokio/src/lib.rs | 199 ++++++++++++++++++++++++ crates/openwire/src/client.rs | 75 ++++----- crates/openwire/src/connection/pool.rs | 42 ++--- crates/openwire/src/policy/follow_up.rs | 101 ++++++++++-- docs/ARCHITECTURE.md | 7 + 5 files changed, 335 insertions(+), 89 deletions(-) diff --git a/crates/openwire-tokio/src/lib.rs b/crates/openwire-tokio/src/lib.rs index bde0822..bb3b9ac 100644 --- a/crates/openwire-tokio/src/lib.rs +++ b/crates/openwire-tokio/src/lib.rs @@ -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 { + inner: R, + positive_ttl: Duration, + negative_ttl: Duration, + cache: std::sync::Arc>>, +} + +#[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), + Err(String), +} + +impl CachingDnsResolver { + /// 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 CachingDnsResolver { + 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 DnsResolver for CachingDnsResolver +where + R: DnsResolver + Clone, +{ + fn resolve( + &self, + ctx: CallContext, + host: String, + port: u16, + ) -> BoxFuture, 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; @@ -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, + addr: SocketAddr, + } + + impl DnsResolver for CountingResolver { + fn resolve( + &self, + ctx: CallContext, + host: String, + port: u16, + ) -> BoxFuture, 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); + } } diff --git a/crates/openwire/src/client.rs b/crates/openwire/src/client.rs index 7d74e1e..2bab3a6 100644 --- a/crates/openwire/src/client.rs +++ b/crates/openwire/src/client.rs @@ -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; @@ -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()), @@ -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); + pub(crate) fn attach_request_admission( - response: Response, + mut response: Response, permit: RequestAdmissionPermit, ) -> Response { - 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, + mut response: Response, ctx: CallContext, state: Arc, ) -> Response { + let admission = response + .extensions_mut() + .remove::() + .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()), ) } @@ -1338,15 +1341,23 @@ struct CallLifecycleBody { inner: Option, ctx: CallContext, state: Arc, + /// Keeps request admission capacity reserved until the body completes. + _admission: Option>, finished: bool, } impl CallLifecycleBody { - fn new(inner: ResponseBody, ctx: CallContext, state: Arc) -> Self { + fn new( + inner: ResponseBody, + ctx: CallContext, + state: Arc, + admission: Option>, + ) -> Self { Self { inner: Some(inner), ctx, state, + _admission: admission, finished: false, } } @@ -1440,34 +1451,6 @@ impl Body for CallLifecycleBody { } } -pin_project! { - struct RequestAdmissionBody { - #[pin] - inner: ResponseBody, - _permit: Option, - } -} - -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, 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}; diff --git a/crates/openwire/src/connection/pool.rs b/crates/openwire/src/connection/pool.rs index 257dbb2..5bcc955 100644 --- a/crates/openwire/src/connection/pool.rs +++ b/crates/openwire/src/connection/pool.rs @@ -196,11 +196,7 @@ impl ConnectionPool { let mut removed = Vec::new(); for candidate_address in &addresses_to_prune { let mut state = lock_mutex(self.shard(candidate_address)); - removed.extend(prune_address( - &self.settings, - &mut state, - candidate_address, - )); + removed.extend(prune_address(&self.settings, &mut state, candidate_address)); } self.finish_removals(removed); @@ -334,18 +330,19 @@ impl ConnectionPool { let mut state = lock_mutex(self.shard(address)); let removed = prune_address(&self.settings, &mut state, address); let stats = match state.by_address.get(address) { - Some(connections) => connections.iter().fold( - PoolStats::default(), - |mut stats, connection| { - stats.total += 1; - match connection.snapshot().allocation { - ConnectionAllocationState::Idle => stats.idle += 1, - ConnectionAllocationState::InUse { .. } => stats.in_use += 1, - ConnectionAllocationState::Closed => {} - } - stats - }, - ), + Some(connections) => { + connections + .iter() + .fold(PoolStats::default(), |mut stats, connection| { + stats.total += 1; + match connection.snapshot().allocation { + ConnectionAllocationState::Idle => stats.idle += 1, + ConnectionAllocationState::InUse { .. } => stats.in_use += 1, + ConnectionAllocationState::Closed => {} + } + stats + }) + } None => PoolStats::default(), }; drop(state); @@ -798,10 +795,7 @@ mod tests { let connection = make_connection(address.clone(), 10); let connection_id = connection.id(); pool.insert(connection.clone()); - assert_eq!( - lock_mutex(&pool.by_id).get(&connection_id), - Some(&address) - ); + assert_eq!(lock_mutex(&pool.by_id).get(&connection_id), Some(&address)); assert_eq!( pool.stats(&address), @@ -1208,10 +1202,7 @@ mod tests { } ); assert!(!lock_mutex(&pool.by_id).contains_key(&stale_id)); - assert_eq!( - lock_mutex(&pool.by_id).get(&live_id), - Some(&live_address) - ); + assert_eq!(lock_mutex(&pool.by_id).get(&live_id), Some(&live_address)); } #[test] @@ -1368,5 +1359,4 @@ mod tests { assert!(pool.acquire(&a).is_none()); assert!(pool.acquire(&b).is_none()); } - } diff --git a/crates/openwire/src/policy/follow_up.rs b/crates/openwire/src/policy/follow_up.rs index 8c4b373..3b4db35 100644 --- a/crates/openwire/src/policy/follow_up.rs +++ b/crates/openwire/src/policy/follow_up.rs @@ -94,7 +94,8 @@ impl Service for FollowUpPolicyService { policy_trace.auth_count = auths; request.extensions_mut().insert(policy_trace); - let snapshot = RequestSnapshot::capture(&request); + let snapshot = + RequestSnapshot::capture(&request, SnapshotCaptureMode::for_policy(&config)); apply_request_cookies(&mut request, config.cookie_jar.as_deref())?; let exchange = Exchange::new(request, ctx.clone(), attempt); let result = if first_network_attempt { @@ -385,8 +386,8 @@ async fn authenticate_response( snapshot.method.clone(), snapshot.uri.clone(), snapshot.version, - (*snapshot.headers).clone(), - (*snapshot.extensions).clone(), + snapshot.headers()?.clone(), + snapshot.extensions()?.clone(), snapshot.body.as_ref().and_then(RequestBody::try_clone), ), AuthResponseState::new(response.status(), response.headers().clone()), @@ -508,14 +509,42 @@ fn store_response_cookies( Ok(()) } +/// Whether this attempt may need a rebuilt request after the network hop. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum SnapshotCaptureMode { + /// Headers/extensions are not needed (no redirect/auth/retry possible). + Light, + /// Full rebuild material may be required after the response. + Rebuildable, +} + +impl SnapshotCaptureMode { + fn for_policy(config: &PolicyConfig) -> Self { + let may_redirect = config + .redirect + .default_policy() + .map(|policy| policy.follow_redirects()) + .unwrap_or(true); + let may_retry = + config.retry.default_config().max_retries() > 0 || config.retry.has_custom_policy(); + let may_auth = + config.auth.authenticator.is_some() || config.auth.proxy_authenticator.is_some(); + if may_redirect || may_retry || may_auth { + Self::Rebuildable + } else { + Self::Light + } + } +} + struct RequestSnapshot { method: Method, uri: Uri, version: Version, - /// Shared so auth / retry rebuilds do not re-clone the full map on each hop. - headers: Arc, - /// Shared request extensions captured at attempt start. - extensions: Arc, + /// Present only when capture mode is rebuildable. + headers: Option>, + /// Present only when capture mode is rebuildable. + extensions: Option>, body: Option, } @@ -526,13 +555,20 @@ fn request_url(uri: &Uri) -> Result { } impl RequestSnapshot { - fn capture(request: &Request) -> Self { + fn capture(request: &Request, mode: SnapshotCaptureMode) -> Self { + let (headers, extensions) = match mode { + SnapshotCaptureMode::Light => (None, None), + SnapshotCaptureMode::Rebuildable => ( + Some(Arc::new(request.headers().clone())), + Some(Arc::new(request.extensions().clone())), + ), + }; Self { method: request.method().clone(), uri: request.uri().clone(), version: request.version(), - headers: Arc::new(request.headers().clone()), - extensions: Arc::new(request.extensions().clone()), + headers, + extensions, body: request.body().try_clone(), } } @@ -541,6 +577,24 @@ impl RequestSnapshot { self.body.is_some() } + fn headers(&self) -> Result<&HeaderMap, WireError> { + self.headers.as_deref().ok_or_else(|| { + WireError::internal( + "request snapshot is missing headers for follow-up rebuild", + std::io::Error::other("light snapshot used for rebuild"), + ) + }) + } + + fn extensions(&self) -> Result<&http::Extensions, WireError> { + self.extensions.as_deref().ok_or_else(|| { + WireError::internal( + "request snapshot is missing extensions for follow-up rebuild", + std::io::Error::other("light snapshot used for rebuild"), + ) + }) + } + fn to_retry_request( &self, policy_trace: PolicyTraceContext, @@ -560,8 +614,8 @@ impl RequestSnapshot { .uri(self.uri.clone()) .version(self.version) .body(body)?; - *request.headers_mut() = (*self.headers).clone(); - *request.extensions_mut() = (*self.extensions).clone(); + *request.headers_mut() = self.headers()?.clone(); + *request.extensions_mut() = self.extensions()?.clone(); let sticky_proxy = request.extensions().get::().cloned(); reset_network_attempt_extensions(request.extensions_mut(), sticky_proxy); request.extensions_mut().insert(policy_trace); @@ -606,7 +660,13 @@ impl RequestSnapshot { self.method }; - let mut headers = Arc::try_unwrap(self.headers).unwrap_or_else(|arc| (*arc).clone()); + let headers = self.headers.ok_or_else(|| { + WireError::internal( + "request snapshot is missing headers for redirect rebuild", + std::io::Error::other("light snapshot used for rebuild"), + ) + })?; + let mut headers = Arc::try_unwrap(headers).unwrap_or_else(|arc| (*arc).clone()); headers.remove(HOST); if !same_origin { strip_sensitive_cross_origin_headers(&mut headers); @@ -619,6 +679,13 @@ impl RequestSnapshot { headers.remove(EXPECT); } + let extensions = self.extensions.ok_or_else(|| { + WireError::internal( + "request snapshot is missing extensions for redirect rebuild", + std::io::Error::other("light snapshot used for rebuild"), + ) + })?; + let mut request = Request::builder() .method(method) .uri(next_uri) @@ -626,7 +693,7 @@ impl RequestSnapshot { .body(body)?; *request.headers_mut() = headers; *request.extensions_mut() = - Arc::try_unwrap(self.extensions).unwrap_or_else(|arc| (*arc).clone()); + Arc::try_unwrap(extensions).unwrap_or_else(|arc| (*arc).clone()); reset_network_attempt_extensions(request.extensions_mut(), selected_proxy); request.extensions_mut().insert(policy_trace); Ok(Some(request)) @@ -797,7 +864,7 @@ mod tests { HeaderValue::from_str(value).expect("header value"), ); } - RequestSnapshot::capture(&request) + RequestSnapshot::capture(&request, super::SnapshotCaptureMode::Rebuildable) } #[test] @@ -939,7 +1006,7 @@ mod tests { Result, >())) .expect("request"); - let snapshot = RequestSnapshot::capture(&request); + let snapshot = RequestSnapshot::capture(&request, super::SnapshotCaptureMode::Rebuildable); let next = snapshot .into_redirect_request( @@ -967,7 +1034,7 @@ mod tests { .header(EXPECT, "100-continue") .body(RequestBody::from_static(b"{}")) .expect("request"); - let snapshot = RequestSnapshot::capture(&request); + let snapshot = RequestSnapshot::capture(&request, super::SnapshotCaptureMode::Rebuildable); let next = snapshot .into_redirect_request( StatusCode::FOUND, diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 3687d8b..1322395 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -358,3 +358,10 @@ cargo test -p openwire --test live_network -- --ignored --test-threads=1 - `ResponseBody::text()` reclaims the collected `Bytes` buffer via `into()` when unique, avoiding an extra `to_vec()` copy on the common path. +- Default clients wrap the system resolver in `CachingDnsResolver` (30s positive / + 5s negative TTL) to avoid repeated system lookups under connection churn. +- Follow-up snapshots stay light when redirects, retries, and authenticators are + all disabled, so the common single-shot path skips header/extension cloning. +- Request admission permits are held via response extensions into the call + lifecycle body, avoiding an extra `BoxBody` layer on the returned response. +