use ahash::AHashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::time::Instant;
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::ConnectionRecorder;
use crate::util::BackoffPolicy;
pub const DEFAULT_MAX_IDLE: Duration = Duration::from_secs(9 * 60);
const RECONNECT_BACKOFF: BackoffPolicy = BackoffPolicy {
initial_backoff: Duration::from_millis(50),
max_backoff: Duration::from_secs(1),
backoff_multiplier: 2.0,
jitter_factor: 0.2,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum ConnectionPurpose {
Data,
Coordination,
}
type PoolKey = (String, ConnectionPurpose);
type DialWaiters = Vec<oneshot::Sender<Result<Arc<BrokerConnection>>>>;
#[derive(Debug)]
struct ReconnectState {
failures: u32,
next_attempt_at: Instant,
last_error: KrafkaError,
}
impl ReconnectState {
fn error_in_window(&self, address: &str, now: Instant) -> KrafkaError {
if !self.last_error.is_retriable() {
return self.last_error.clone();
}
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionRefused,
format!(
"not reconnecting to {address} for another {:?} after {} failed attempt(s); \
last error: {}",
self.next_attempt_at.saturating_duration_since(now),
self.failures,
self.last_error
),
))
}
}
struct PoolState {
by_id: AHashMap<BrokerId, Arc<BrokerConnection>>,
by_key: AHashMap<PoolKey, Arc<BrokerConnection>>,
dialing: AHashMap<PoolKey, DialWaiters>,
reconnect: AHashMap<String, ReconnectState>,
generation: u64,
fallback_warned: bool,
}
impl PoolState {
fn new() -> Self {
Self {
by_id: AHashMap::new(),
by_key: AHashMap::new(),
dialing: AHashMap::new(),
reconnect: AHashMap::new(),
generation: 0,
fallback_warned: false,
}
}
fn counted_connections(&self) -> usize {
self.by_key.values().filter(|c| c.is_alive()).count() + self.dialing.len()
}
}
enum DialStart {
Ready(Arc<BrokerConnection>),
Wait(oneshot::Receiver<Result<Arc<BrokerConnection>>>),
CapReached(usize),
}
pub struct ConnectionPool {
state: Arc<RwLock<PoolState>>,
config: ConnectionConfig,
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: Arc::new(RwLock::new(PoolState::new())),
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(crate) fn recorder(&self) -> &Arc<ConnectionRecorder> {
self.config.connection_metrics()
}
pub fn metrics(&self) -> crate::metrics::ConnectionMetrics {
self.recorder().snapshot()
}
pub(crate) fn open_connections(&self) -> Vec<Arc<BrokerConnection>> {
self.state
.read()
.by_key
.iter()
.filter(|((_, purpose), conn)| *purpose == ConnectionPurpose::Data && conn.is_usable())
.map(|(_, conn)| Arc::clone(conn))
.collect()
}
#[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
}
pub async fn get_connection(&self, address: &str) -> Result<Arc<BrokerConnection>> {
self.get(address, ConnectionPurpose::Data).await
}
pub async fn get_coordinator_connection(&self, address: &str) -> Result<Arc<BrokerConnection>> {
self.get(address, ConnectionPurpose::Coordination).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(address, ConnectionPurpose::Data).await?;
self.state.write().by_id.insert(broker_id, conn.clone());
Ok(conn)
}
async fn get(
&self,
address: &str,
purpose: ConnectionPurpose,
) -> Result<Arc<BrokerConnection>> {
let key = (address.to_string(), purpose);
{
let s = self.state.read();
if let Some(conn) = s.by_key.get(&key)
&& conn.is_usable()
{
return Ok(conn.clone());
}
}
let rx = match self.start_dial(key)? {
DialStart::Ready(conn) => return Ok(conn),
DialStart::Wait(rx) => rx,
DialStart::CapReached(limit) => {
if purpose == ConnectionPurpose::Coordination {
self.record_coordination_fallback(address, limit);
return Box::pin(self.get(address, ConnectionPurpose::Data)).await;
}
return Err(connection_cap_error(limit, address));
}
};
rx.await.map_err(|_| {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
format!("connection attempt to {address} was abandoned"),
))
})?
}
fn start_dial(&self, key: PoolKey) -> Result<DialStart> {
let now = Instant::now();
let mut s = self.state.write();
if let Some(conn) = s.by_key.get(&key)
&& conn.is_usable()
{
return Ok(DialStart::Ready(conn.clone()));
}
if let Some(waiters) = s.dialing.get_mut(&key) {
let (tx, rx) = oneshot::channel();
waiters.push(tx);
return Ok(DialStart::Wait(rx));
}
if let Some(state) = s.reconnect.get(&key.0)
&& now < state.next_attempt_at
{
return Err(state.error_in_window(&key.0, now));
}
if let Some(stale) = s.by_key.remove(&key) {
s.by_id.retain(|_, c| !Arc::ptr_eq(c, &stale));
if stale.is_alive() {
info!(
address = %key.0,
purpose = ?key.1,
"Replacing connection due to SASL session expiry (KIP-368)"
);
stale.close_when_idle();
}
}
if let Some(limit) = self.max_total_connections
&& s.counted_connections() >= limit
{
return Ok(DialStart::CapReached(limit));
}
let (tx, rx) = oneshot::channel();
s.dialing.insert(key.clone(), vec![tx]);
let generation = s.generation;
drop(s);
self.spawn_dial(key, generation);
Ok(DialStart::Wait(rx))
}
fn spawn_dial(&self, key: PoolKey, generation: u64) {
let state = Arc::clone(&self.state);
let config = self.config.clone();
tokio::spawn(async move {
let connect_timeout = config.connect_timeout;
let address = key.0.clone();
let result = match tokio::time::timeout(
connect_timeout,
BrokerConnection::connect(&address, config),
)
.await
{
Ok(Ok(conn)) => Ok(Arc::new(conn)),
Ok(Err(e)) => Err(e),
Err(_) => Err(KrafkaError::timeout(format!(
"connection to {address} was not established within {connect_timeout:?}"
))),
};
let (waiters, result) = {
let mut s = state.write();
let waiters = s.dialing.remove(&key).unwrap_or_default();
let result = match result {
Ok(conn) if s.generation != generation => {
conn.close_when_idle();
Err(pool_closed_error(&address))
}
Ok(conn) => {
s.reconnect.remove(&address);
s.by_key.insert(key, conn.clone());
Ok(conn)
}
Err(e) => {
let failures = s.reconnect.get(&address).map_or(0, |r| r.failures) + 1;
let backoff = RECONNECT_BACKOFF.calculate_backoff(failures);
warn!(
address = %address,
failures,
backoff_ms = backoff.as_millis() as u64,
error = %e,
"Connection attempt failed"
);
s.reconnect.insert(
address,
ReconnectState {
failures,
next_attempt_at: Instant::now() + backoff,
last_error: e.clone(),
},
);
Err(e)
}
};
(waiters, result)
};
for waiter in waiters {
let _ = waiter.send(result.clone());
}
});
}
fn record_coordination_fallback(&self, address: &str, limit: usize) {
self.recorder().record_coordination_fallback();
let first = {
let mut s = self.state.write();
!std::mem::replace(&mut s.fallback_warned, true)
};
if first {
warn!(
address = %address,
max_total_connections = limit,
"Connection cap reached: coordination requests share the data connection; \
heartbeats may wait behind fetches. Raise the cap to isolate them."
);
}
}
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_keys: Vec<PoolKey> = s
.by_key
.iter()
.filter(|(_, c)| c.idle_duration() >= max_idle)
.map(|(key, _)| key.clone())
.collect();
for key in stale_keys {
if let Some(c) = s.by_key.remove(&key) {
if c.idle_duration() >= max_idle {
removed.push(c);
} else {
s.by_key.insert(key, 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"
);
for conn in removed {
conn.close_now();
}
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;
};
provider.bind_metrics(Arc::clone(self.recorder()));
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();
}
}
#[allow(clippy::unused_async)]
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 (connections, dialing) = {
let mut s = self.state.write();
s.generation += 1;
s.reconnect.clear();
let mut connections: Vec<_> = s.by_id.drain().map(|(_, c)| c).collect();
connections.extend(s.by_key.drain().map(|(_, c)| c));
(connections, std::mem::take(&mut s.dialing))
};
for ((address, _), waiters) in dialing {
let err = pool_closed_error(&address);
for waiter in waiters {
let _ = waiter.send(Err(err.clone()));
}
}
for conn in connections {
conn.close_now();
}
}
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
}
}
fn connection_cap_error(limit: usize, address: &str) -> KrafkaError {
KrafkaError::network(std::io::Error::other(format!(
"connection pool limit reached: {limit} connections open, none to {address} \
(raise `max_connections` or reduce the number of brokers in use)"
)))
}
fn pool_closed_error(address: &str) -> KrafkaError {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionAborted,
format!("pool closed while connecting to {address}"),
))
}
#[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;
}
#[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_key.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_key
.insert(("b1:9092".to_string(), ConnectionPurpose::Data), stale);
}
assert_eq!(pool.evict_idle(), 1);
{
let s = pool.state.read();
assert!(s.by_id.is_empty());
assert!(s.by_key.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"
);
}
#[tokio::test]
async fn token_fetches_are_reported_to_the_pools_metrics() {
let config = oauth_pool_config();
let pool = Arc::new(ConnectionPool::new(config));
pool.start_token_refresh();
let provider = pool
.config
.auth
.as_ref()
.and_then(|a| a.oauthbearer_provider())
.expect("the config carries a provider")
.clone();
provider.provide_token().await.expect("provider succeeds");
assert_eq!(
pool.metrics().oauth_token_fetches,
1,
"the connection-path fetch must reach the pool's metrics"
);
pool.close_all().await;
}
#[test]
fn metrics_are_bound_even_when_the_refresh_task_cannot_start() {
let pool = Arc::new(ConnectionPool::new(oauth_pool_config()));
pool.start_token_refresh();
let provider = pool
.config
.auth
.as_ref()
.and_then(|a| a.oauthbearer_provider())
.expect("the config carries a provider")
.clone();
let metrics = Arc::clone(pool.recorder());
tokio::runtime::Builder::new_current_thread()
.build()
.expect("runtime")
.block_on(async { provider.provide_token().await })
.expect("provider succeeds");
assert_eq!(metrics.oauth_token_fetches.get(), 1);
}
#[cfg(feature = "test-broker")]
#[tokio::test]
async fn a_session_expired_connection_is_replaced_at_the_cap() {
let broker = crate::testing::FakeBroker::start().await.unwrap();
let addr = broker.bootstrap_servers();
let pool = ConnectionPool::new(ConnectionConfig::default()).with_max_total_connections(1);
let expired = Arc::new(BrokerConnection::test_stub_session_expired(&addr));
assert!(expired.is_alive() && !expired.is_usable());
pool.state
.write()
.by_key
.insert((addr.clone(), ConnectionPurpose::Data), expired.clone());
let fresh = pool
.get_connection(&addr)
.await
.expect("replaced, not refused");
assert!(!Arc::ptr_eq(&expired, &fresh));
assert!(fresh.is_usable());
}
#[cfg(feature = "test-broker")]
#[tokio::test]
async fn a_successful_dial_resets_the_backoff() {
let broker = crate::testing::FakeBroker::start().await.unwrap();
let addr = broker.bootstrap_servers();
let pool = ConnectionPool::new(ConnectionConfig::default());
pool.state.write().reconnect.insert(
addr.clone(),
ReconnectState {
failures: 4,
next_attempt_at: Instant::now(),
last_error: KrafkaError::timeout("earlier failure"),
},
);
pool.get_connection(&addr).await.unwrap();
assert!(!pool.state.read().reconnect.contains_key(&addr));
}
#[tokio::test]
async fn the_backoff_window_repeats_a_non_retriable_error() {
let pool = ConnectionPool::new(ConnectionConfig::default());
pool.state.write().reconnect.insert(
"b1:9092".to_string(),
ReconnectState {
failures: 1,
next_attempt_at: Instant::now() + Duration::from_secs(60),
last_error: KrafkaError::auth("bad credentials"),
},
);
let err = pool
.get_connection("b1:9092")
.await
.err()
.expect("fails inside the window");
assert!(matches!(err, KrafkaError::Auth { .. }), "{err:?}");
}
#[tokio::test(start_paused = true)]
async fn tls_reload_fires_at_its_interval() {
use crate::auth::{AuthConfig, TlsConfig};
let ca = format!("{}/src/auth/testdata/ca.pem", env!("CARGO_MANIFEST_DIR"));
let config = ConnectionConfig::builder()
.auth(AuthConfig::ssl(TlsConfig::new().with_ca_cert(ca)))
.build()
.unwrap();
let pool = Arc::new(ConnectionPool::new(config));
let interval = Duration::from_secs(60);
let loaded = |pool: &ConnectionPool| pool.config.tls_connector.load().is_some();
let start = Instant::now();
pool.start_tls_reload(interval);
tokio::time::sleep_until(start + interval - Duration::from_millis(1)).await;
assert!(!loaded(&pool), "reloaded before the interval");
tokio::time::sleep_until(start + interval).await;
while !loaded(&pool) {
tokio::time::sleep(Duration::from_millis(1)).await;
}
assert!(
start.elapsed() <= interval + Duration::from_millis(5),
"reloaded {:?} after start",
start.elapsed()
);
}
}