use ahash::AHashMap;
use std::sync::Arc;
use std::time::Duration;
use parking_lot::{Mutex, RwLock};
use tokio::sync::oneshot;
use tokio::task::JoinHandle;
use tracing::{debug, info, warn};
use super::connection::{BrokerConnection, ConnectionConfig};
use crate::BrokerId;
use crate::error::{KrafkaError, Result};
use crate::metrics::ConnectionMetrics;
use crate::util::BackoffPolicy;
#[derive(Debug, Clone)]
pub struct ConnectionRetryConfig {
pub(crate) max_retries: u32,
pub(crate) backoff: BackoffPolicy,
}
impl Default for ConnectionRetryConfig {
fn default() -> Self {
Self {
max_retries: 3,
backoff: BackoffPolicy {
jitter_factor: 0.2, ..BackoffPolicy::default()
},
}
}
}
impl ConnectionRetryConfig {
pub fn builder() -> ConnectionRetryConfigBuilder {
ConnectionRetryConfigBuilder::default()
}
#[inline]
pub fn max_retries(&self) -> u32 {
self.max_retries
}
#[inline]
pub fn initial_backoff(&self) -> Duration {
self.backoff.initial_backoff
}
#[inline]
pub fn max_backoff(&self) -> Duration {
self.backoff.max_backoff
}
#[inline]
pub fn backoff_multiplier(&self) -> f64 {
self.backoff.backoff_multiplier
}
#[inline]
pub fn jitter_factor(&self) -> f64 {
self.backoff.jitter_factor
}
#[inline]
fn calculate_backoff(&self, attempt: u32) -> Duration {
self.backoff.calculate_backoff(attempt)
}
}
#[must_use = "builders do nothing until .build() is called"]
#[derive(Debug, Default)]
pub struct ConnectionRetryConfigBuilder {
config: ConnectionRetryConfig,
}
impl ConnectionRetryConfigBuilder {
pub fn max_retries(mut self, retries: u32) -> Self {
self.config.max_retries = retries;
self
}
pub fn initial_backoff(mut self, duration: Duration) -> Self {
self.config.backoff.initial_backoff = duration;
self
}
pub fn max_backoff(mut self, duration: Duration) -> Self {
self.config.backoff.max_backoff = duration;
self
}
pub fn backoff_multiplier(mut self, multiplier: f64) -> Self {
self.config.backoff.backoff_multiplier = if multiplier.is_finite() && multiplier > 0.0 {
multiplier
} else {
1.0
};
self
}
pub fn jitter_factor(mut self, factor: f64) -> Self {
self.config.backoff.jitter_factor = if factor.is_finite() {
factor.clamp(0.0, 1.0)
} else {
0.0
};
self
}
pub fn build(self) -> ConnectionRetryConfig {
self.config
}
}
pub const DEFAULT_MAX_IDLE: Duration = Duration::from_secs(9 * 60);
const CLOSE_ALL_TIMEOUT: Duration = Duration::from_secs(10);
type ConnectingWaiters = AHashMap<String, Vec<oneshot::Sender<Result<Arc<BrokerConnection>>>>>;
struct ReconnectGuard {
connecting: Arc<Mutex<ConnectingWaiters>>,
address: Option<String>,
}
impl ReconnectGuard {
fn new(connecting: &Arc<Mutex<ConnectingWaiters>>, address: String) -> Self {
Self {
connecting: Arc::clone(connecting),
address: Some(address),
}
}
fn defuse(&mut self) {
self.address = None;
}
}
impl Drop for ReconnectGuard {
fn drop(&mut self) {
let Some(address) = self.address.take() else {
return;
};
let mut guard = self.connecting.lock();
let waiters = guard.remove(&address).unwrap_or_default();
let err = KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
format!("reconnection to {address} was cancelled"),
));
for waiter in waiters {
let _ = waiter.send(Err(err.clone()));
}
}
}
struct PoolState {
by_id: AHashMap<BrokerId, Arc<BrokerConnection>>,
by_addr: AHashMap<String, Arc<BrokerConnection>>,
}
impl PoolState {
fn new() -> Self {
Self {
by_id: AHashMap::new(),
by_addr: AHashMap::new(),
}
}
}
pub struct ConnectionPool {
state: RwLock<PoolState>,
connecting: Arc<Mutex<ConnectingWaiters>>,
config: ConnectionConfig,
retry_config: ConnectionRetryConfig,
max_idle: Option<Duration>,
max_total_connections: Option<usize>,
evictor_handle: Mutex<Option<JoinHandle<()>>>,
oauth_refresh_handle: Mutex<Option<JoinHandle<()>>>,
tls_reload_handle: Mutex<Option<JoinHandle<()>>>,
}
impl ConnectionPool {
pub fn new(config: ConnectionConfig) -> Self {
Self {
state: RwLock::new(PoolState::new()),
connecting: Arc::new(Mutex::new(AHashMap::new())),
config,
retry_config: ConnectionRetryConfig::default(),
max_idle: Some(DEFAULT_MAX_IDLE),
max_total_connections: None,
evictor_handle: Mutex::new(None),
oauth_refresh_handle: Mutex::new(None),
tls_reload_handle: Mutex::new(None),
}
}
pub fn with_retry_config(
config: ConnectionConfig,
retry_config: ConnectionRetryConfig,
) -> Self {
Self {
state: RwLock::new(PoolState::new()),
connecting: Arc::new(Mutex::new(AHashMap::new())),
config,
retry_config,
max_idle: Some(DEFAULT_MAX_IDLE),
max_total_connections: None,
evictor_handle: Mutex::new(None),
oauth_refresh_handle: Mutex::new(None),
tls_reload_handle: Mutex::new(None),
}
}
pub fn start(config: ConnectionConfig) -> Arc<Self> {
let pool = Arc::new(Self::new(config));
pool.start_idle_evictor();
pool
}
#[inline]
pub fn metrics(&self) -> Arc<ConnectionMetrics> {
self.config.connection_metrics()
}
#[must_use]
pub fn with_max_idle(mut self, max_idle: Option<Duration>) -> Self {
self.max_idle = max_idle;
self
}
#[inline]
pub fn max_idle(&self) -> Option<Duration> {
self.max_idle
}
#[must_use]
pub fn with_max_total_connections(mut self, limit: impl Into<Option<usize>>) -> Self {
self.max_total_connections = limit.into();
self
}
#[inline]
pub fn max_total_connections(&self) -> Option<usize> {
self.max_total_connections
}
pub async fn refresh_tls(&self) -> crate::error::Result<()> {
self.config.refresh_tls().await
}
fn connect_attempt_budget(&self) -> Duration {
self.config
.connect_timeout
.saturating_mul(3)
.max(Duration::from_secs(1))
}
fn total_reconnect_budget(&self) -> Duration {
let attempts = self.retry_config.max_retries.saturating_add(1);
let attempt_cost = self.connect_attempt_budget().saturating_mul(attempts);
let backoff_cost = self
.retry_config
.max_backoff()
.saturating_mul(self.retry_config.max_retries);
attempt_cost.saturating_add(backoff_cost)
}
async fn reconnect_with_backoff(&self, address: &str) -> Result<Arc<BrokerConnection>> {
let mut last_error: Option<KrafkaError> = None;
let overall_deadline = tokio::time::Instant::now() + self.total_reconnect_budget();
for attempt in 0..=self.retry_config.max_retries {
if tokio::time::Instant::now() >= overall_deadline {
warn!(
address = %address,
attempt = attempt,
"Reconnect budget exhausted; giving up"
);
break;
}
if attempt > 0 {
let backoff = self.retry_config.calculate_backoff(attempt);
debug!(
address = %address,
attempt = attempt,
max_retries = self.retry_config.max_retries,
backoff_ms = backoff.as_millis(),
"Retrying connection after backoff"
);
tokio::time::sleep(backoff).await;
}
let attempt_deadline =
(tokio::time::Instant::now() + self.connect_attempt_budget()).min(overall_deadline);
let attempt_result = match tokio::time::timeout_at(
attempt_deadline,
BrokerConnection::connect(address, self.config.clone()),
)
.await
{
Ok(r) => r,
Err(_) => Err(KrafkaError::timeout(format!(
"connection to {address} did not complete within {:?} \
(TCP established but handshake stalled?)",
self.connect_attempt_budget()
))),
};
match attempt_result {
Ok(conn) => {
if attempt > 0 {
info!(
address = %address,
attempt = attempt,
"Successfully reconnected after retries"
);
}
return Ok(Arc::new(conn));
}
Err(e) => {
if !e.is_retriable() {
warn!(
address = %address,
error = %e,
"Non-retriable connection error, not retrying"
);
return Err(e);
}
warn!(
address = %address,
attempt = attempt,
max_retries = self.retry_config.max_retries,
error = %e,
"Connection attempt failed"
);
last_error = Some(e);
}
}
}
Err(last_error.unwrap_or_else(|| {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionRefused,
format!(
"Failed to connect to {} after {} retries",
address, self.retry_config.max_retries
),
))
}))
}
async fn get_or_reconnect(&self, address: &str) -> Result<Arc<BrokerConnection>> {
{
let s = self.state.read();
if s.by_addr
.get(address)
.is_some_and(|c| c.is_alive() && c.needs_reauthentication())
{
info!(
address = %address,
"Replacing connection due to SASL session expiry (KIP-368)"
);
}
}
enum CoalesceAction {
AlreadyConnected(Arc<BrokerConnection>),
WaitForPeer(oneshot::Receiver<Result<Arc<BrokerConnection>>>),
Reconnect(String),
}
let existing = {
let s = self.state.read();
s.by_addr.get(address).filter(|c| c.is_usable()).cloned()
};
let action = {
let mut connecting = self.connecting.lock();
if let Some(conn) = existing {
CoalesceAction::AlreadyConnected(conn)
} else if let Some(waiters) = connecting.get_mut(address) {
let (tx, rx) = oneshot::channel();
waiters.push(tx);
CoalesceAction::WaitForPeer(rx)
} else {
let addr_owned = address.to_string();
connecting.insert(addr_owned.clone(), Vec::new());
CoalesceAction::Reconnect(addr_owned)
}
};
let addr_owned = match action {
CoalesceAction::AlreadyConnected(conn) => return Ok(conn),
CoalesceAction::WaitForPeer(rx) => {
let waiter_budget = self.total_reconnect_budget() + Duration::from_secs(1);
return tokio::time::timeout(waiter_budget, rx)
.await
.map_err(|_| {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"timed out after {waiter_budget:?} waiting for an in-flight \
reconnection to {address}"
),
))
})?
.map_err(|_| {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
format!("reconnection to {address} was cancelled"),
))
})?;
}
CoalesceAction::Reconnect(addr_owned) => addr_owned,
};
let mut guard = ReconnectGuard::new(&self.connecting, addr_owned.clone());
if let Some(limit) = self.max_total_connections {
let current = self.state.read().by_addr.len();
if current >= limit {
let err = KrafkaError::config(format!(
"connection pool limit reached: {current}/{limit} connections open \
(address={address}); raise `max_total_connections` or reduce broker count"
));
let waiters = self
.connecting
.lock()
.remove(&addr_owned)
.unwrap_or_default();
for waiter in waiters {
let _ = waiter.send(Err(err.clone()));
}
guard.defuse();
return Err(err);
}
}
let result = self.reconnect_with_backoff(address).await;
let waiters = self.connecting.lock().remove(address).unwrap_or_default();
let final_result = match result {
Ok(conn) => {
let mut s = self.state.write();
if let Some(limit) = self.max_total_connections {
if s.by_addr.len() >= limit {
drop(s);
let overflow = conn.clone();
tokio::spawn(async move { overflow.close().await });
Err(KrafkaError::config(format!(
"connection pool limit reached: {limit} connections open \
(address={addr_owned}); raise `max_total_connections` or reduce broker count"
)))
} else {
s.by_addr.insert(addr_owned, conn.clone());
Ok(conn)
}
} else {
s.by_addr.insert(addr_owned, conn.clone());
Ok(conn)
}
}
Err(e) => Err(e),
};
for waiter in waiters {
let _ = waiter.send(final_result.clone());
}
guard.defuse();
final_result
}
pub async fn get_connection(&self, address: &str) -> Result<Arc<BrokerConnection>> {
{
let s = self.state.read();
if let Some(conn) = s.by_addr.get(address)
&& conn.is_usable()
{
return Ok(conn.clone());
}
}
self.get_or_reconnect(address).await
}
pub async fn get_connection_by_id(
&self,
broker_id: BrokerId,
address: &str,
) -> Result<Arc<BrokerConnection>> {
{
let s = self.state.read();
if let Some(conn) = s.by_id.get(&broker_id)
&& conn.is_usable()
&& conn.address() == address
{
return Ok(conn.clone());
}
}
let conn = self.get_or_reconnect(address).await?;
{
let mut s = self.state.write();
let keep_existing = s
.by_id
.get(&broker_id)
.is_some_and(|c| c.is_usable() && c.address() == address);
if !keep_existing {
s.by_id.insert(broker_id, conn.clone());
}
}
Ok(conn)
}
pub fn evict_idle(&self) -> usize {
let Some(max_idle) = self.max_idle else {
return 0;
};
let mut removed: Vec<Arc<BrokerConnection>> = Vec::new();
{
let mut s = self.state.write();
let stale_ids: Vec<BrokerId> = s
.by_id
.iter()
.filter(|(_, c)| c.idle_duration() >= max_idle)
.map(|(id, _)| *id)
.collect();
for id in stale_ids {
if let Some(c) = s.by_id.remove(&id) {
if c.idle_duration() >= max_idle {
removed.push(c);
} else {
s.by_id.insert(id, c);
}
}
}
let stale_addrs: Vec<String> = s
.by_addr
.iter()
.filter(|(_, c)| c.idle_duration() >= max_idle)
.map(|(addr, _)| addr.clone())
.collect();
for addr in stale_addrs {
if let Some(c) = s.by_addr.remove(&addr) {
if c.idle_duration() >= max_idle {
removed.push(c);
} else {
s.by_addr.insert(addr, c);
}
}
}
}
if removed.is_empty() {
return 0;
}
removed.sort_by_key(|c| Arc::as_ptr(c) as usize);
removed.dedup_by(|a, b| Arc::ptr_eq(a, b));
let count = removed.len();
debug!(
evicted = count,
max_idle_ms = max_idle.as_millis(),
"Evicted idle connections"
);
if tokio::runtime::Handle::try_current().is_ok() {
for conn in removed {
tokio::spawn(async move { conn.close().await });
}
}
count
}
pub fn start_idle_evictor(self: &Arc<Self>) {
self.start_token_refresh();
let Some(max_idle) = self.max_idle else {
return;
};
if tokio::runtime::Handle::try_current().is_err() {
warn!("start_idle_evictor called outside a Tokio runtime; idle eviction disabled");
return;
}
let interval = (max_idle / 9)
.max(Duration::from_secs(1))
.min(Duration::from_secs(60));
let weak = Arc::downgrade(self);
let handle = tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
ticker.tick().await;
loop {
ticker.tick().await;
let Some(pool) = weak.upgrade() else {
break;
};
pool.evict_idle();
}
});
if let Some(prev) = self.evictor_handle.lock().replace(handle) {
prev.abort();
}
}
pub fn start_token_refresh(self: &Arc<Self>) {
let Some(provider) = self
.config
.auth
.as_ref()
.and_then(|a| a.oauthbearer_provider())
else {
return;
};
if tokio::runtime::Handle::try_current().is_err() {
warn!(
"start_token_refresh called outside a Tokio runtime; OAUTHBEARER \
proactive refresh disabled. Tokens are still refreshed lazily on \
the connection path."
);
return;
}
let refresh_handle = provider.start_refresh_task();
if let Some(prev) = self.oauth_refresh_handle.lock().replace(refresh_handle) {
prev.abort();
}
}
pub fn start_tls_reload(self: &Arc<Self>, interval: Duration) {
if interval.is_zero() {
warn!("start_tls_reload called with a zero interval; ignoring");
return;
}
if self
.config
.auth
.as_ref()
.and_then(|a| a.tls_config.as_ref())
.is_none()
{
return;
}
if tokio::runtime::Handle::try_current().is_err() {
warn!(
"start_tls_reload called outside a Tokio runtime; automatic TLS \
reloading disabled. Call `refresh_tls()` explicitly instead."
);
return;
}
let weak = Arc::downgrade(self);
let handle = tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
ticker.tick().await; loop {
ticker.tick().await;
let Some(pool) = weak.upgrade() else {
break;
};
match pool.refresh_tls().await {
Ok(()) => debug!("Periodic TLS reload completed"),
Err(e) => warn!(
error = %e,
"Periodic TLS reload failed; keeping the previously loaded \
certificates and retrying on the next tick"
),
}
}
});
if let Some(prev) = self.tls_reload_handle.lock().replace(handle) {
prev.abort();
}
}
pub async fn close_all(&self) {
if let Some(handle) = self.evictor_handle.lock().take() {
handle.abort();
}
if let Some(handle) = self.oauth_refresh_handle.lock().take() {
handle.abort();
}
if let Some(handle) = self.tls_reload_handle.lock().take() {
handle.abort();
}
let (by_id, by_addr) = {
let mut s = self.state.write();
(
s.by_id.drain().map(|(_, c)| c).collect::<Vec<_>>(),
s.by_addr.drain().map(|(_, c)| c).collect::<Vec<_>>(),
)
};
let mut seen = AHashMap::with_capacity(by_id.len() + by_addr.len());
for conn in by_id.into_iter().chain(by_addr) {
seen.entry(Arc::as_ptr(&conn) as usize).or_insert(conn);
}
{
let mut connecting = self.connecting.lock();
for (addr, waiters) in connecting.drain() {
let err = KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
format!("pool closed while reconnecting to {addr}"),
));
for waiter in waiters {
let _ = waiter.send(Err(err.clone()));
}
}
}
let closes = seen
.into_values()
.map(|conn| async move { conn.close().await });
if tokio::time::timeout(CLOSE_ALL_TIMEOUT, futures::future::join_all(closes))
.await
.is_err()
{
warn!(
"close_all timed out after {CLOSE_ALL_TIMEOUT:?}; remaining sockets \
will be torn down when their last Arc drops"
);
}
}
pub fn len(&self) -> usize {
let s = self.state.read();
s.by_id.values().filter(|c| c.is_usable()).count()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn test_connection_pool_new() {
let pool = ConnectionPool::new(ConnectionConfig::default());
let _ = pool;
}
#[test]
fn test_connection_retry_config_default() {
let config = ConnectionRetryConfig::default();
assert_eq!(config.max_retries, 3);
assert_eq!(config.initial_backoff(), Duration::from_millis(100));
assert_eq!(config.max_backoff(), Duration::from_secs(10));
assert_eq!(config.backoff_multiplier(), 2.0);
}
#[test]
fn test_calculate_backoff() {
let config = ConnectionRetryConfig::builder().jitter_factor(0.0).build();
assert_eq!(config.calculate_backoff(0), Duration::ZERO);
assert_eq!(config.calculate_backoff(1), Duration::from_millis(100));
assert_eq!(config.calculate_backoff(2), Duration::from_millis(200));
assert_eq!(config.calculate_backoff(3), Duration::from_millis(400));
}
#[test]
fn test_calculate_backoff_capped() {
let config = ConnectionRetryConfig::builder()
.max_retries(10)
.initial_backoff(Duration::from_secs(1))
.max_backoff(Duration::from_secs(5))
.backoff_multiplier(10.0)
.jitter_factor(0.0)
.build();
assert_eq!(config.calculate_backoff(2), Duration::from_secs(5));
}
#[test]
fn test_calculate_backoff_handles_max_attempt() {
let config = ConnectionRetryConfig::builder()
.max_retries(u32::MAX)
.jitter_factor(0.0)
.build();
assert_eq!(config.calculate_backoff(u32::MAX), config.max_backoff());
}
#[test]
fn test_connection_pool_with_retry_config() {
let retry_config = ConnectionRetryConfig::builder()
.max_retries(5)
.initial_backoff(Duration::from_millis(50))
.max_backoff(Duration::from_secs(5))
.backoff_multiplier(3.0)
.jitter_factor(0.2)
.build();
let pool = ConnectionPool::with_retry_config(ConnectionConfig::default(), retry_config);
assert_eq!(pool.retry_config.max_retries, 5);
assert_eq!(
pool.retry_config.initial_backoff(),
Duration::from_millis(50)
);
}
#[tokio::test]
async fn test_pool_close_all_clears_both_maps() {
let pool = ConnectionPool::new(ConnectionConfig::default());
{
let s = pool.state.read();
assert!(s.by_id.is_empty());
assert!(s.by_addr.is_empty());
}
pool.close_all().await;
}
#[test]
fn test_max_idle_default_matches_java_client() {
let pool = ConnectionPool::new(ConnectionConfig::default());
assert_eq!(pool.max_idle(), Some(Duration::from_millis(9 * 60 * 1000)));
assert_eq!(DEFAULT_MAX_IDLE, Duration::from_secs(540));
}
#[test]
fn test_with_max_idle_none_disables_eviction() {
let pool = ConnectionPool::new(ConnectionConfig::default()).with_max_idle(None);
assert_eq!(pool.max_idle(), None);
assert_eq!(pool.evict_idle(), 0);
}
#[test]
fn test_evict_idle_on_empty_pool_is_noop() {
let pool = ConnectionPool::new(ConnectionConfig::default());
assert_eq!(pool.evict_idle(), 0);
}
#[tokio::test]
async fn test_start_idle_evictor_installs_and_aborts_task() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
assert!(pool.evictor_handle.lock().is_none());
pool.start_idle_evictor();
assert!(pool.evictor_handle.lock().is_some());
pool.start_idle_evictor();
assert!(pool.evictor_handle.lock().is_some());
pool.close_all().await;
assert!(pool.evictor_handle.lock().is_none());
}
#[tokio::test]
async fn test_start_idle_evictor_noop_when_max_idle_disabled() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()).with_max_idle(None));
pool.start_idle_evictor();
assert!(pool.evictor_handle.lock().is_none());
}
#[test]
fn test_start_idle_evictor_noop_outside_tokio_runtime() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
pool.start_idle_evictor();
assert!(
pool.evictor_handle.lock().is_none(),
"evictor must not be installed without a Tokio runtime"
);
}
#[test]
fn test_evict_idle_removes_stale_from_both_maps() {
let pool = ConnectionPool::new(ConnectionConfig::default())
.with_max_idle(Some(Duration::from_millis(100)));
let stale = Arc::new(BrokerConnection::test_stub_idle_for(
"b1:9092",
Duration::from_secs(10),
));
{
let mut s = pool.state.write();
s.by_id.insert(1, stale.clone());
s.by_addr.insert("b1:9092".to_string(), stale);
}
assert_eq!(pool.evict_idle(), 1);
{
let s = pool.state.read();
assert!(s.by_id.is_empty());
assert!(s.by_addr.is_empty());
}
}
#[test]
fn test_evict_idle_retains_fresh_and_evicts_stale() {
let pool = ConnectionPool::new(ConnectionConfig::default())
.with_max_idle(Some(Duration::from_millis(100)));
let stale = Arc::new(BrokerConnection::test_stub_idle_for(
"b1:9092",
Duration::from_secs(10),
));
let fresh = Arc::new(BrokerConnection::test_stub_idle_for(
"b2:9092",
Duration::from_millis(10),
));
{
let mut s = pool.state.write();
s.by_id.insert(1, stale);
s.by_id.insert(2, fresh);
}
assert_eq!(pool.evict_idle(), 1);
let s = pool.state.read();
assert!(!s.by_id.contains_key(&1));
assert!(s.by_id.contains_key(&2));
}
#[test]
fn test_evict_idle_rescued_after_refresh() {
let pool = ConnectionPool::new(ConnectionConfig::default())
.with_max_idle(Some(Duration::from_millis(100)));
let conn = Arc::new(BrokerConnection::test_stub_idle_for(
"b1:9092",
Duration::from_secs(10),
));
conn.test_mark_fresh();
pool.state.write().by_id.insert(1, conn);
assert_eq!(pool.evict_idle(), 0);
assert!(pool.state.read().by_id.contains_key(&1));
}
#[test]
fn test_max_total_connections_default_is_none() {
let pool = ConnectionPool::new(ConnectionConfig::default());
assert_eq!(pool.max_total_connections(), None);
}
#[test]
fn test_with_max_total_connections_sets_limit() {
let pool =
ConnectionPool::new(ConnectionConfig::default()).with_max_total_connections(10usize);
assert_eq!(pool.max_total_connections(), Some(10));
}
#[test]
fn test_with_max_total_connections_none_removes_limit() {
let pool = ConnectionPool::new(ConnectionConfig::default())
.with_max_total_connections(5usize)
.with_max_total_connections(None);
assert_eq!(pool.max_total_connections(), None);
}
fn oauth_pool_config() -> ConnectionConfig {
ConnectionConfig::builder()
.auth(crate::auth::AuthConfig::sasl_oauthbearer_provider(
|| async { Ok(crate::auth::OAuthBearerToken::new("jwt")) },
))
.build()
.unwrap()
}
#[tokio::test]
async fn test_token_refresh_starts_even_when_idle_eviction_disabled() {
let pool = Arc::new(ConnectionPool::new(oauth_pool_config()).with_max_idle(None));
pool.start_idle_evictor();
assert!(
pool.evictor_handle.lock().is_none(),
"eviction is disabled, as configured"
);
assert!(
pool.oauth_refresh_handle.lock().is_some(),
"token refresh must run regardless of max_idle"
);
pool.close_all().await;
}
#[tokio::test]
async fn test_token_refresh_starts_alongside_idle_evictor() {
let pool = Arc::new(ConnectionPool::new(oauth_pool_config()));
pool.start_idle_evictor();
assert!(pool.evictor_handle.lock().is_some());
assert!(pool.oauth_refresh_handle.lock().is_some());
pool.close_all().await;
assert!(
pool.oauth_refresh_handle.lock().is_none(),
"aborted on close"
);
}
#[tokio::test]
async fn test_start_token_refresh_is_noop_without_provider() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
pool.start_token_refresh();
assert!(pool.oauth_refresh_handle.lock().is_none());
}
#[test]
fn test_start_token_refresh_is_noop_outside_runtime() {
let pool = Arc::new(ConnectionPool::new(oauth_pool_config()));
pool.start_token_refresh();
assert!(
pool.oauth_refresh_handle.lock().is_none(),
"must not panic in tokio::spawn without a runtime"
);
}
#[test]
fn test_reconnect_budgets_are_finite_and_ordered() {
let pool = ConnectionPool::new(ConnectionConfig::default());
let attempt = pool.connect_attempt_budget();
let total = pool.total_reconnect_budget();
assert!(attempt >= pool.config.connect_timeout);
assert!(
total > attempt,
"the total budget must cover every attempt plus backoff"
);
assert!(total < Duration::from_secs(600), "budget must stay sane");
}
}