use ahash::AHashMap;
use futures::FutureExt;
use std::sync::Arc;
use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::time::{Duration, Instant, SystemTime};
use arc_swap::ArcSwap;
use bytes::{Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
#[cfg(feature = "socks5")]
use tokio::net::TcpSocket;
use tokio::sync::{mpsc, oneshot};
use tokio::time::timeout;
#[cfg(feature = "socks5")]
use tokio::time::timeout_at;
use tokio_rustls::TlsConnector;
use tokio_util::time::{DelayQueue, delay_queue};
use tracing::{debug, error, info, trace, warn};
use crate::CorrelationId;
use crate::auth::msk_iam::MAX_SIGV4_CLOCK_SKEW_SECS;
use crate::auth::tls::build_tls_connector;
use crate::auth::{
AuthConfig, ChannelBinding, SaslMechanism, SecurityProtocol, connect_tls,
extract_tls_server_end_point,
};
use crate::error::{ErrorCode, KrafkaError, ProtocolErrorKind, Result};
use crate::metrics::ConnectionMetrics;
struct ConnectionLoopParams {
address: String,
high_priority_rx: mpsc::Receiver<ConnectionCommand>,
normal_priority_rx: mpsc::Receiver<ConnectionCommand>,
request_timeout: Duration,
stats: Arc<ConnectionStats>,
metrics: Arc<ConnectionMetrics>,
max_response_size: usize,
max_in_flight_requests: usize,
max_high_priority_bypasses: usize,
}
#[cfg(feature = "socks5")]
#[derive(Clone)]
pub struct ProxyConfig {
address: String,
credentials: Option<ProxyCredentials>,
}
#[cfg(feature = "socks5")]
impl ProxyConfig {
pub fn new(address: impl Into<String>) -> Self {
Self {
address: address.into(),
credentials: None,
}
}
pub fn with_credentials(
address: impl Into<String>,
username: impl Into<String>,
password: impl Into<String>,
) -> Self {
Self {
address: address.into(),
credentials: Some(ProxyCredentials {
username: zeroize::Zeroizing::new(username.into()),
password: zeroize::Zeroizing::new(password.into()),
}),
}
}
#[inline]
pub fn address(&self) -> &str {
&self.address
}
#[inline]
pub fn credentials(&self) -> Option<&ProxyCredentials> {
self.credentials.as_ref()
}
}
#[cfg(feature = "socks5")]
impl std::fmt::Debug for ProxyConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProxyConfig")
.field("address", &self.address)
.field(
"credentials",
if self.credentials.is_some() {
&"[REDACTED]"
} else {
&"None"
},
)
.finish()
}
}
#[cfg(feature = "socks5")]
#[derive(Clone, zeroize::ZeroizeOnDrop)]
pub struct ProxyCredentials {
username: zeroize::Zeroizing<String>,
password: zeroize::Zeroizing<String>,
}
#[cfg(feature = "socks5")]
impl ProxyCredentials {
#[inline]
pub fn username(&self) -> &str {
&self.username
}
#[inline]
pub fn password(&self) -> &str {
&self.password
}
}
#[cfg(feature = "socks5")]
impl std::fmt::Debug for ProxyCredentials {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ProxyCredentials")
.field("username", &self.username.as_str())
.field("password", &"[REDACTED]")
.finish()
}
}
use crate::protocol::{
ApiKey, ApiVersionRange, ApiVersionsRequest, ApiVersionsResponse, Decoder, Encoder,
FinalizedFeature, RequestHeader, ResponseHeader, SaslAuthenticateRequest,
SaslAuthenticateResponse, SaslHandshakeRequest, SaslHandshakeResponse, SupportedFeature,
};
use crate::util::{CorrelationIdGenerator, NO_RESPONSE_CORRELATION_ID, extract_sni_hostname};
use super::secure::{ChallengeResponse, SaslAuthenticator};
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RequestPriority {
High,
Normal,
}
impl RequestPriority {
#[inline]
pub fn for_api_key(api_key: ApiKey) -> Self {
match api_key {
ApiKey::Heartbeat
| ApiKey::ConsumerGroupHeartbeat
| ApiKey::ShareGroupHeartbeat
| ApiKey::JoinGroup
| ApiKey::SyncGroup
| ApiKey::LeaveGroup
| ApiKey::OffsetCommit => Self::High,
ApiKey::Metadata => Self::High,
ApiKey::FindCoordinator => Self::High,
ApiKey::LeaderAndIsr => Self::High,
ApiKey::ApiVersions => Self::High,
_ => Self::Normal,
}
}
}
pub const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Clone)]
pub struct ConnectionConfig {
pub(crate) connect_timeout: Duration,
pub(crate) request_timeout: Duration,
pub(crate) send_buffer_size: Option<usize>,
pub(crate) recv_buffer_size: Option<usize>,
pub(crate) nodelay: bool,
pub(crate) client_id: String,
pub(crate) high_priority_channel_capacity: usize,
pub(crate) normal_priority_channel_capacity: usize,
pub(crate) max_response_size: usize,
pub(crate) max_in_flight_requests: usize,
pub(crate) max_high_priority_bypasses_per_round: usize,
pub(crate) auth: Option<AuthConfig>,
pub(crate) tls_connector: Arc<ArcSwap<Option<TlsConnector>>>,
pub(crate) tcp_keepalive: Option<Duration>,
pub(crate) connection_attempt_delay: Duration,
pub(crate) msk_iam_clock_offset_secs: Arc<AtomicI64>,
pub(crate) connection_metrics: Arc<ConnectionMetrics>,
#[cfg(feature = "socks5")]
pub(crate) proxy: Option<ProxyConfig>,
}
impl std::fmt::Debug for ConnectionConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut s = f.debug_struct("ConnectionConfig");
s.field("connect_timeout", &self.connect_timeout)
.field("request_timeout", &self.request_timeout)
.field("send_buffer_size", &self.send_buffer_size)
.field("recv_buffer_size", &self.recv_buffer_size)
.field("nodelay", &self.nodelay)
.field("client_id", &self.client_id)
.field(
"high_priority_channel_capacity",
&self.high_priority_channel_capacity,
)
.field(
"normal_priority_channel_capacity",
&self.normal_priority_channel_capacity,
)
.field("max_response_size", &self.max_response_size)
.field("max_in_flight_requests", &self.max_in_flight_requests)
.field(
"max_high_priority_bypasses_per_round",
&self.max_high_priority_bypasses_per_round,
)
.field("auth", &self.auth)
.field("tls_connector", &self.tls_connector.load().is_some())
.field("tcp_keepalive", &self.tcp_keepalive)
.field("connection_attempt_delay", &self.connection_attempt_delay)
.field(
"msk_iam_clock_offset_secs",
&self.msk_iam_clock_offset_secs.load(Ordering::Relaxed),
);
#[cfg(feature = "socks5")]
s.field("proxy", &self.proxy);
s.finish()
}
}
impl Default for ConnectionConfig {
fn default() -> Self {
#[allow(clippy::expect_used)]
ConnectionConfigBuilder::default()
.build()
.expect("default ConnectionConfig values are always valid")
}
}
impl ConnectionConfig {
pub fn builder() -> ConnectionConfigBuilder {
ConnectionConfigBuilder::default()
}
pub async fn init_tls(&mut self) -> Result<()> {
if let Some(ref auth) = self.auth
&& let Some(ref tls_config) = auth.tls_config
{
let connector = build_tls_connector(tls_config).await?;
self.tls_connector.store(Arc::new(Some(connector)));
}
Ok(())
}
pub async fn refresh_tls(&self) -> Result<()> {
if let Some(ref auth) = self.auth
&& let Some(ref tls_config) = auth.tls_config
{
let connector = build_tls_connector(tls_config).await?;
self.tls_connector.store(Arc::new(Some(connector)));
info!("TLS connector refreshed from disk");
}
Ok(())
}
#[inline]
pub fn connect_timeout(&self) -> Duration {
self.connect_timeout
}
#[inline]
pub fn request_timeout(&self) -> Duration {
self.request_timeout
}
#[inline]
pub fn send_buffer_size(&self) -> Option<usize> {
self.send_buffer_size
}
#[inline]
pub fn recv_buffer_size(&self) -> Option<usize> {
self.recv_buffer_size
}
#[inline]
pub fn nodelay(&self) -> bool {
self.nodelay
}
#[inline]
pub fn client_id(&self) -> &str {
&self.client_id
}
#[inline]
pub fn high_priority_channel_capacity(&self) -> usize {
self.high_priority_channel_capacity
}
#[inline]
pub fn normal_priority_channel_capacity(&self) -> usize {
self.normal_priority_channel_capacity
}
#[inline]
pub fn max_response_size(&self) -> usize {
self.max_response_size
}
#[inline]
pub fn max_in_flight_requests(&self) -> usize {
self.max_in_flight_requests
}
#[inline]
pub fn auth(&self) -> Option<&AuthConfig> {
self.auth.as_ref()
}
#[inline]
pub fn connection_attempt_delay(&self) -> Duration {
self.connection_attempt_delay
}
#[inline]
pub fn connection_metrics(&self) -> Arc<ConnectionMetrics> {
self.connection_metrics.clone()
}
#[cfg(feature = "socks5")]
#[inline]
pub fn proxy(&self) -> Option<&ProxyConfig> {
self.proxy.as_ref()
}
}
#[must_use = "builders do nothing until .build() is called"]
#[derive(Debug)]
pub struct ConnectionConfigBuilder(ConnectionConfig);
impl Default for ConnectionConfigBuilder {
fn default() -> Self {
ConnectionConfigBuilder(ConnectionConfig {
connect_timeout: DEFAULT_CONNECT_TIMEOUT,
request_timeout: Duration::from_secs(30),
send_buffer_size: None,
recv_buffer_size: None,
nodelay: true,
client_id: "krafka".to_string(),
high_priority_channel_capacity: 64,
normal_priority_channel_capacity: 256,
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 10,
max_high_priority_bypasses_per_round: 4,
auth: None,
tls_connector: Arc::new(ArcSwap::new(Arc::new(None))),
tcp_keepalive: Some(Duration::from_secs(60)),
connection_attempt_delay: Duration::from_millis(250),
msk_iam_clock_offset_secs: Arc::new(AtomicI64::new(0)),
connection_metrics: Arc::new(ConnectionMetrics::default()),
#[cfg(feature = "socks5")]
proxy: None,
})
}
}
impl ConnectionConfigBuilder {
pub fn connect_timeout(mut self, timeout: Duration) -> Self {
self.0.connect_timeout = timeout;
self
}
pub fn request_timeout(mut self, timeout: Duration) -> Self {
self.0.request_timeout = timeout;
self
}
pub fn client_id(mut self, client_id: impl Into<String>) -> Self {
self.0.client_id = client_id.into();
self
}
pub fn nodelay(mut self, nodelay: bool) -> Self {
self.0.nodelay = nodelay;
self
}
pub fn high_priority_channel_capacity(mut self, capacity: usize) -> Self {
self.0.high_priority_channel_capacity = capacity.max(16);
self
}
pub fn normal_priority_channel_capacity(mut self, capacity: usize) -> Self {
self.0.normal_priority_channel_capacity = capacity.max(64);
self
}
pub fn max_response_size(mut self, size: usize) -> Self {
self.0.max_response_size = size.max(1024); self
}
pub fn socket_send_buffer(mut self, bytes: Option<usize>) -> Self {
self.0.send_buffer_size = bytes;
self
}
pub fn socket_receive_buffer(mut self, bytes: Option<usize>) -> Self {
self.0.recv_buffer_size = bytes;
self
}
pub fn max_in_flight_requests(mut self, max: usize) -> Self {
self.0.max_in_flight_requests = max.max(1);
self
}
pub fn max_high_priority_bypasses_per_round(mut self, n: usize) -> Self {
self.0.max_high_priority_bypasses_per_round = n.max(1);
self
}
pub fn auth(mut self, auth: AuthConfig) -> Self {
self.0.auth = Some(auth);
self
}
pub fn tcp_keepalive(mut self, interval: Option<Duration>) -> Self {
self.0.tcp_keepalive = interval;
self
}
pub fn connection_attempt_delay(mut self, delay: Duration) -> Self {
self.0.connection_attempt_delay = delay;
self
}
pub fn connection_metrics(mut self, metrics: Arc<ConnectionMetrics>) -> Self {
self.0.connection_metrics = metrics;
self
}
#[cfg(feature = "socks5")]
pub fn proxy(mut self, proxy: ProxyConfig) -> Self {
self.0.proxy = Some(proxy);
self
}
pub fn build(self) -> crate::error::Result<ConnectionConfig> {
const MAX_CLIENT_ID_LEN: usize = i16::MAX as usize;
if self.0.client_id.len() > MAX_CLIENT_ID_LEN {
return Err(crate::error::KrafkaError::config(format!(
"client_id is {} bytes, exceeding the Kafka wire limit of {MAX_CLIENT_ID_LEN}",
self.0.client_id.len()
)));
}
if self.0.request_timeout < self.0.connect_timeout {
return Err(crate::error::KrafkaError::config(format!(
"request_timeout ({:?}) must be >= connect_timeout ({:?}); \
otherwise all requests time out before the connection completes. \
Lower connect_timeout to match if you want a shorter request_timeout",
self.0.request_timeout, self.0.connect_timeout
)));
}
const WARN_CEILING_BYTES: usize = 1024 * 1024 * 1024; let ceiling = self
.0
.max_response_size
.saturating_mul(self.0.max_in_flight_requests);
if ceiling > WARN_CEILING_BYTES {
tracing::warn!(
max_response_size = self.0.max_response_size,
max_in_flight_requests = self.0.max_in_flight_requests,
ceiling_bytes = ceiling,
"ConnectionConfig memory ceiling ({} × {} = {} bytes) exceeds 1 GiB; \
consider lowering max_response_size or max_in_flight_requests",
self.0.max_response_size,
self.0.max_in_flight_requests,
ceiling,
);
}
Ok(self.0)
}
}
pub(crate) const MAX_SASL_FRAME_BYTES: usize = 64 * 1024;
const HANDSHAKE_READ_CHUNK: usize = 8 * 1024;
const MAX_HONOURED_THROTTLE_MS: i32 = 5 * 60 * 1000;
fn leading_throttle_time_ms(api_key: ApiKey, api_version: i16, body: &[u8]) -> Option<i32> {
if api_version < api_key.leading_throttle_time_min_version()? {
return None;
}
let bytes: [u8; 4] = body.get(..4)?.try_into().ok()?;
let throttle_time_ms = i32::from_be_bytes(bytes);
if throttle_time_ms > 0 && throttle_time_ms <= MAX_HONOURED_THROTTLE_MS {
Some(throttle_time_ms)
} else {
None
}
}
const CLOSE_SEND_TIMEOUT: Duration = Duration::from_secs(5);
fn connection_closed_error() -> KrafkaError {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"connection closed",
))
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct BrokerFeatures {
pub negotiated_version: i16,
pub supported: Vec<SupportedFeature>,
pub finalized_epoch: i64,
pub finalized: Vec<FinalizedFeature>,
}
impl BrokerFeatures {
#[must_use]
pub fn finalized_level(&self, name: &str) -> Option<i16> {
if self.finalized_epoch < 0 {
return None;
}
self.finalized
.iter()
.find(|f| f.name == name)
.map(|f| f.max_version_level)
}
#[must_use]
pub fn supported_range(&self, name: &str) -> Option<(i16, i16)> {
self.supported
.iter()
.find(|f| f.name == name)
.map(|f| (f.min_version, f.max_version))
}
}
struct PendingRequest {
response_tx: oneshot::Sender<Result<Bytes>>,
api_key: ApiKey,
api_version: i16,
_permit: tokio::sync::OwnedSemaphorePermit,
}
enum ConnectionCommand {
Request {
data: Bytes,
correlation_id: CorrelationId,
api_key: ApiKey,
api_version: i16,
response_tx: oneshot::Sender<Result<Bytes>>,
timeout: Duration,
permit: tokio::sync::OwnedSemaphorePermit,
},
FireAndForget { data: Bytes },
Close,
}
pub struct BrokerConnection {
address: String,
config: ConnectionConfig,
correlation_id_gen: Arc<CorrelationIdGenerator>,
high_priority_tx: mpsc::Sender<ConnectionCommand>,
normal_priority_tx: mpsc::Sender<ConnectionCommand>,
api_versions: Arc<parking_lot::Mutex<AHashMap<ApiKey, ApiVersionRange>>>,
broker_features: Arc<parking_lot::Mutex<BrokerFeatures>>,
alive: Arc<std::sync::atomic::AtomicBool>,
session_expiry: Option<Instant>,
stats: Arc<ConnectionStats>,
throttle_until: Arc<parking_lot::Mutex<Instant>>,
created_at: Instant,
last_used_nanos: AtomicU64,
in_flight: Arc<tokio::sync::Semaphore>,
}
#[derive(Debug, Default)]
#[non_exhaustive]
pub struct ConnectionStats {
pub high_priority_requests: AtomicU64,
pub normal_priority_requests: AtomicU64,
pub high_priority_bypasses: AtomicU64,
pub high_priority_bypass_yields: AtomicU64,
}
impl ConnectionStats {
#[inline]
pub fn high_priority_count(&self) -> u64 {
self.high_priority_requests.load(Ordering::Relaxed)
}
#[inline]
pub fn normal_priority_count(&self) -> u64 {
self.normal_priority_requests.load(Ordering::Relaxed)
}
#[inline]
pub fn bypass_count(&self) -> u64 {
self.high_priority_bypasses.load(Ordering::Relaxed)
}
#[inline]
pub fn bypass_yield_count(&self) -> u64 {
self.high_priority_bypass_yields.load(Ordering::Relaxed)
}
}
impl BrokerConnection {
pub async fn connect(address: &str, config: ConnectionConfig) -> Result<Self> {
let stream = Self::establish_tcp(address, &config).await?;
stream.set_nodelay(config.nodelay)?;
debug!("Connected to broker at {address}");
let (high_priority_tx, high_priority_rx) =
mpsc::channel(config.high_priority_channel_capacity);
let (normal_priority_tx, normal_priority_rx) =
mpsc::channel(config.normal_priority_channel_capacity);
let alive = Arc::new(std::sync::atomic::AtomicBool::new(true));
let alive_clone = alive.clone();
let stats = Arc::new(ConnectionStats::default());
let stats_clone = stats.clone();
let mut connection = Self {
address: address.to_string(),
config: config.clone(),
correlation_id_gen: Arc::new(CorrelationIdGenerator::new()),
high_priority_tx,
normal_priority_tx,
api_versions: Arc::new(parking_lot::Mutex::new(AHashMap::new())),
broker_features: Arc::new(parking_lot::Mutex::new(BrokerFeatures::default())),
alive,
session_expiry: None,
stats,
throttle_until: Arc::new(parking_lot::Mutex::new(Instant::now())),
created_at: Instant::now(),
last_used_nanos: AtomicU64::new(0),
in_flight: Arc::new(tokio::sync::Semaphore::new(config.max_in_flight_requests)),
};
let request_timeout = config.request_timeout;
let handshake_deadline = tokio::time::Instant::now() + config.connect_timeout;
let loop_params = ConnectionLoopParams {
address: address.to_string(),
high_priority_rx,
normal_priority_rx,
request_timeout,
stats: stats_clone,
metrics: config.connection_metrics.clone(),
max_response_size: config.max_response_size,
max_in_flight_requests: config.max_in_flight_requests,
max_high_priority_bypasses: config.max_high_priority_bypasses_per_round,
};
if let Some(auth) = config.auth.as_ref().filter(|a| a.requires_tls()) {
let tls_config = auth
.tls_config
.as_ref()
.ok_or_else(|| KrafkaError::config("TLS required but no TLS config provided"))?;
let connector = match &**config.tls_connector.load() {
Some(c) => c.clone(),
None => build_tls_connector(tls_config).await?,
};
let hostname = extract_sni_hostname(address)?;
let tls_start = std::time::Instant::now();
let tls_stream = tokio::time::timeout_at(
handshake_deadline,
connect_tls(
stream,
hostname,
tls_config.sni_hostname.as_deref(),
&connector,
),
)
.await
.map_err(|_| {
KrafkaError::timeout(format!(
"TLS handshake with {address} did not complete within {:?}",
config.connect_timeout
))
})??;
config
.connection_metrics
.record_tls_handshake(tls_start.elapsed());
info!("TLS handshake completed for {address}");
if auth.requires_sasl() {
let mut tls_stream = tls_stream;
let channel_binding = if auth.scram_channel_binding {
extract_tls_server_end_point(&tls_stream)
.map(ChannelBinding::TlsServerEndPoint)
.unwrap_or(ChannelBinding::None)
} else {
debug!(
"SCRAM channel binding disabled by configuration for {address}; \
using unbound n,, GS2 framing"
);
ChannelBinding::None
};
let session_lifetime_ms = Self::perform_sasl_handshake(
&mut tls_stream,
auth,
address,
&config.client_id,
request_timeout,
handshake_deadline,
channel_binding,
&config.msk_iam_clock_offset_secs,
)
.await?;
connection.session_expiry =
Self::effective_session_expiry(session_lifetime_ms, auth);
let (reader, writer) = tokio::io::split(tls_stream);
config.connection_metrics.record_connect();
Self::spawn_connection_task(reader, writer, loop_params, alive_clone);
} else {
let (reader, writer) = tokio::io::split(tls_stream);
config.connection_metrics.record_connect();
Self::spawn_connection_task(reader, writer, loop_params, alive_clone);
}
} else if let Some(auth) = config.auth.as_ref().filter(|a| a.requires_sasl()) {
let mut stream = stream;
let session_lifetime_ms = Self::perform_sasl_handshake(
&mut stream,
auth,
address,
&config.client_id,
request_timeout,
handshake_deadline,
ChannelBinding::None,
&config.msk_iam_clock_offset_secs,
)
.await?;
connection.session_expiry = Self::effective_session_expiry(session_lifetime_ms, auth);
let (reader, writer) = stream.into_split();
config.connection_metrics.record_connect();
Self::spawn_connection_task(reader, writer, loop_params, alive_clone);
} else {
let (reader, writer) = stream.into_split();
config.connection_metrics.record_connect();
Self::spawn_connection_task(reader, writer, loop_params, alive_clone);
}
connection.fetch_api_versions().await?;
Ok(connection)
}
async fn establish_tcp(
address: &str,
config: &ConnectionConfig,
) -> Result<tokio::net::TcpStream> {
#[cfg(feature = "socks5")]
if let Some(ref proxy) = config.proxy {
return Self::connect_via_proxy(address, proxy, config).await;
}
Self::connect_direct(address, config).await
}
async fn connect_direct(
address: &str,
config: &ConnectionConfig,
) -> Result<tokio::net::TcpStream> {
super::happy_eyeballs::connect_happy_eyeballs(address, config).await
}
#[cfg(feature = "socks5")]
async fn connect_via_proxy(
address: &str,
proxy: &ProxyConfig,
config: &ConnectionConfig,
) -> Result<tokio::net::TcpStream> {
use tokio_socks::tcp::Socks5Stream;
debug!("Connecting to {address} via SOCKS5 proxy {}", proxy.address);
let deadline = tokio::time::Instant::now() + config.connect_timeout;
let addrs: Vec<std::net::SocketAddr> =
timeout_at(deadline, tokio::net::lookup_host(&proxy.address))
.await
.map_err(|_| KrafkaError::timeout("SOCKS5 proxy DNS resolution"))?
.map_err(KrafkaError::network)?
.collect();
if addrs.is_empty() {
return Err(KrafkaError::invalid_state(format!(
"no addresses resolved for SOCKS5 proxy '{}'",
proxy.address
)));
}
let proxy_addr = addrs[0];
let socket = Self::create_socket(proxy_addr, config)?;
let proxy_stream = timeout_at(deadline, async {
let tcp = socket
.connect(proxy_addr)
.await
.map_err(KrafkaError::network)?;
let socks = if let Some(ref creds) = proxy.credentials {
Socks5Stream::connect_with_password_and_socket(
tcp,
address,
creds.username(),
creds.password(),
)
.await
} else {
Socks5Stream::connect_with_socket(tcp, address).await
}
.map_err(|e| {
KrafkaError::network(std::io::Error::other(format!("SOCKS5 proxy error: {e}")))
})?;
Ok::<_, KrafkaError>(socks.into_inner())
})
.await
.map_err(|_| KrafkaError::timeout("SOCKS5 proxy connection"))??;
info!(
"SOCKS5 tunnel established to {address} via {}",
proxy.address
);
Ok(proxy_stream)
}
#[cfg(feature = "socks5")]
fn create_socket(addr: std::net::SocketAddr, config: &ConnectionConfig) -> Result<TcpSocket> {
super::happy_eyeballs::create_socket(addr, config)
}
#[allow(clippy::too_many_arguments)]
async fn perform_sasl_handshake<S>(
stream: &mut S,
auth: &AuthConfig,
address: &str,
client_id: &str,
request_timeout: Duration,
deadline: tokio::time::Instant,
channel_binding: ChannelBinding,
msk_iam_clock_offset_secs: &Arc<AtomicI64>,
) -> Result<i64>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let resolved_msk_iam;
let auth = if let Some(resolved) = timeout(request_timeout, auth.resolve_msk_iam_provider())
.await
.map_err(|_| KrafkaError::timeout("MSK IAM credential provider"))??
{
debug!("Resolved MSK IAM credentials from provider for {address}");
resolved_msk_iam = resolved;
&resolved_msk_iam
} else {
auth
};
let resolved_auth;
let auth = if let Some(resolved) =
timeout(request_timeout, auth.resolve_provider_to_token())
.await
.map_err(|_| KrafkaError::timeout("OAUTHBEARER token provider"))??
{
debug!("Resolved OAUTHBEARER token from provider for {address}");
resolved_auth = resolved;
&resolved_auth
} else {
auth
};
let mut authenticator = SaslAuthenticator::new(auth, channel_binding)?
.ok_or_else(|| KrafkaError::auth("Failed to create SASL authenticator"))?;
if auth.security_protocol == SecurityProtocol::SaslPlaintext
&& auth.sasl_mechanism == Some(SaslMechanism::Plain)
{
warn!(
"SASL PLAIN credentials will be sent in cleartext to {}. \
Use SASL_SSL (sasl_plain_ssl) for production environments.",
address
);
}
let hostname = extract_sni_hostname(address)?;
let clock_offset = msk_iam_clock_offset_secs.load(Ordering::Relaxed);
authenticator.set_msk_host(auth, hostname, clock_offset)?;
let mechanism_name = authenticator.mechanism_name().to_string();
debug!("Starting SASL handshake with mechanism {mechanism_name} for {address}");
let handshake_request = SaslHandshakeRequest::new(&mechanism_name);
let mut encoder = Encoder::with_capacity(64);
let pos = encoder.start_message();
let header = RequestHeader::new(ApiKey::SaslHandshake, 1, 0).with_client_id(client_id);
header.encode_v1(encoder.buffer_mut())?;
handshake_request.encode_v1(encoder.buffer_mut())?;
encoder.finish_message(pos)?;
stream
.write_all(&encoder.take())
.await
.map_err(KrafkaError::network)?;
stream.flush().await.map_err(KrafkaError::network)?;
let mut response_buf =
Self::read_handshake_frame(stream, deadline, "SaslHandshake").await?;
let _header = ResponseHeader::decode(&mut response_buf, ApiKey::SaslHandshake, 1)?;
let handshake_response = SaslHandshakeResponse::decode_v0(&mut response_buf)?;
if !handshake_response.is_ok() {
return Err(KrafkaError::auth(format!(
"SASL handshake failed: {:?}. Broker supports: {:?}",
handshake_response.error_code, handshake_response.enabled_mechanisms
)));
}
debug!(
"SASL handshake accepted mechanism {mechanism_name}, broker supports: {:?}",
handshake_response.enabled_mechanisms
);
let initial_bytes = authenticator.initial_response()?;
Self::send_sasl_authenticate(stream, &initial_bytes, client_id).await?;
let auth_response = Self::read_sasl_authenticate_response(stream, deadline).await?;
if !auth_response.error_code.is_ok() {
let err_msg = auth_response.error_message.unwrap_or_default();
if mechanism_name == "AWS_MSK_IAM" {
let lower = err_msg.to_ascii_lowercase();
if lower.contains("signature expired")
|| lower.contains("signature not yet current")
|| lower.contains("request time too")
|| lower.contains("clock")
|| lower.contains("time skew")
{
let skew = Self::extract_clock_skew_secs(&err_msg);
let prev = msk_iam_clock_offset_secs.load(Ordering::Relaxed);
let nudge = if skew != 0 {
skew
} else if lower.contains("expired") || lower.contains("past") {
300
} else {
-300
};
let adjusted =
Self::clamp_msk_iam_clock_offset_secs(prev.saturating_add(nudge));
msk_iam_clock_offset_secs.store(adjusted, Ordering::Relaxed);
warn!(
"MSK IAM auth failed with possible clock skew ({}); \
adjusted clock offset to {}s for next attempt",
err_msg, adjusted,
);
}
}
return Err(KrafkaError::auth(format!(
"SASL authentication failed: {:?} - {}",
auth_response.error_code, err_msg
)));
}
let mut session_lifetime_ms = auth_response.session_lifetime_ms;
const MAX_SASL_ROUNDS: usize = 10;
if !authenticator.is_complete() {
let mut challenge = auth_response.auth_bytes;
let mut rounds = 0;
loop {
match authenticator.process_challenge(&challenge).await? {
ChallengeResponse::Done => break,
ChallengeResponse::AckThenFail { ack, error } => {
let _ = Self::send_sasl_authenticate(stream, &ack, client_id).await;
return Err(error);
}
ChallengeResponse::Continue(response_bytes) => {
rounds += 1;
if rounds > MAX_SASL_ROUNDS {
return Err(KrafkaError::auth(format!(
"SASL challenge-response exceeded {MAX_SASL_ROUNDS} rounds"
)));
}
Self::send_sasl_authenticate(stream, &response_bytes, client_id).await?;
let resp = Self::read_sasl_authenticate_response(stream, deadline).await?;
if !resp.error_code.is_ok() {
return Err(KrafkaError::auth(format!(
"SASL authentication step failed: {:?} - {}",
resp.error_code,
resp.error_message.unwrap_or_default()
)));
}
session_lifetime_ms = resp.session_lifetime_ms;
challenge = resp.auth_bytes;
if authenticator.is_complete() {
break;
}
}
}
}
}
info!("SASL authentication completed ({mechanism_name}) for {address}");
if session_lifetime_ms > 0 {
debug!("Broker reported session lifetime of {session_lifetime_ms}ms for {address}");
}
Ok(session_lifetime_ms)
}
async fn send_sasl_authenticate<S>(
stream: &mut S,
auth_bytes: &[u8],
client_id: &str,
) -> Result<()>
where
S: AsyncWrite + Unpin,
{
let request = SaslAuthenticateRequest::new(auth_bytes.to_vec());
let mut encoder = Encoder::with_capacity(64 + auth_bytes.len());
let pos = encoder.start_message();
let header = RequestHeader::new(ApiKey::SaslAuthenticate, 1, 1).with_client_id(client_id);
header.encode(encoder.buffer_mut())?;
request.encode_v1(encoder.buffer_mut())?;
encoder.finish_message(pos)?;
stream
.write_all(&encoder.take())
.await
.map_err(KrafkaError::network)?;
stream.flush().await.map_err(KrafkaError::network)?;
Ok(())
}
async fn read_sasl_authenticate_response<S>(
stream: &mut S,
deadline: tokio::time::Instant,
) -> Result<SaslAuthenticateResponse>
where
S: AsyncRead + Unpin,
{
let mut buf = Self::read_handshake_frame(stream, deadline, "SaslAuthenticate").await?;
let _header = ResponseHeader::decode(&mut buf, ApiKey::SaslAuthenticate, 1)?;
SaslAuthenticateResponse::decode_v1(&mut buf)
}
async fn read_handshake_frame<S>(
stream: &mut S,
deadline: tokio::time::Instant,
what: &str,
) -> Result<Bytes>
where
S: AsyncRead + Unpin,
{
tokio::time::timeout_at(
deadline,
Self::read_framed_response(stream, MAX_SASL_FRAME_BYTES),
)
.await
.map_err(|_| {
KrafkaError::timeout(format!(
"timed out reading the {what} response during SASL handshake"
))
})?
}
async fn read_framed_response<S>(stream: &mut S, max_len: usize) -> Result<Bytes>
where
S: AsyncRead + Unpin,
{
let mut len_buf = [0u8; 4];
stream
.read_exact(&mut len_buf)
.await
.map_err(KrafkaError::network)?;
let len_i32 = i32::from_be_bytes(len_buf);
if len_i32 <= 0 || (len_i32 as usize) > max_len {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::InvalidLength,
format!(
"Invalid pre-authentication response length: {len_i32} (max: {max_len}); \
refusing to allocate on an unauthenticated peer's say-so"
),
));
}
let len = len_i32 as usize;
let mut body = Vec::with_capacity(len.min(HANDSHAKE_READ_CHUNK));
let mut chunk = [0u8; HANDSHAKE_READ_CHUNK];
while body.len() < len {
let want = (len - body.len()).min(HANDSHAKE_READ_CHUNK);
let n = stream
.read(&mut chunk[..want])
.await
.map_err(KrafkaError::network)?;
if n == 0 {
return Err(KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
format!(
"peer closed during SASL handshake after {} of {len} bytes",
body.len()
),
)));
}
body.extend_from_slice(&chunk[..n]);
}
Ok(Bytes::from(body))
}
fn extract_clock_skew_secs(error_msg: &str) -> i64 {
const AWS_TS_LEN: usize = 16;
let bytes = error_msg.as_bytes();
if bytes.len() < AWS_TS_LEN {
return 0;
}
let mut last_server_unix: Option<i64> = None;
for i in 0..=bytes.len() - AWS_TS_LEN {
if bytes[i + 8] != b'T' || bytes[i + 15] != b'Z' {
continue;
}
if let Some(unix_secs) = Self::parse_aws_ts_unix(&bytes[i..i + AWS_TS_LEN]) {
last_server_unix = Some(unix_secs);
}
}
if let Some(server_unix) = last_server_unix {
let local_unix = SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0);
return server_unix - local_unix;
}
0
}
fn clamp_msk_iam_clock_offset_secs(offset: i64) -> i64 {
offset.clamp(-MAX_SIGV4_CLOCK_SKEW_SECS, MAX_SIGV4_CLOCK_SKEW_SECS)
}
fn parse_aws_ts_unix(s: &[u8]) -> Option<i64> {
debug_assert_eq!(s.len(), 16);
debug_assert_eq!(s[8], b'T');
debug_assert_eq!(s[15], b'Z');
fn d2(hi: u8, lo: u8) -> Option<u32> {
let h = hi.wrapping_sub(b'0');
let l = lo.wrapping_sub(b'0');
if h > 9 || l > 9 {
return None;
}
Some(h as u32 * 10 + l as u32)
}
let year = {
let [a, b, c, d] = [s[0], s[1], s[2], s[3]].map(|x| x.wrapping_sub(b'0'));
if a > 9 || b > 9 || c > 9 || d > 9 {
return None;
}
a as i64 * 1000 + b as i64 * 100 + c as i64 * 10 + d as i64
};
let month = d2(s[4], s[5])? as i64;
let day = d2(s[6], s[7])? as i64;
let hour = d2(s[9], s[10])? as i64;
let min = d2(s[11], s[12])? as i64;
let sec = d2(s[13], s[14])? as i64;
if !(1..=12).contains(&month) {
return None;
}
if hour > 23 || min > 59 || sec > 59 {
return None;
}
let is_leap = (year % 4 == 0 && year % 100 != 0) || (year % 400 == 0);
const DAYS: [i64; 13] = [0, 31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31];
let max_day = if month == 2 && is_leap {
29
} else {
DAYS[month as usize]
};
if !(1..=max_day).contains(&day) {
return None;
}
let a = (14 - month) / 12;
let y = year + 4800 - a;
let m = month + 12 * a - 3;
let jdn = day + (153 * m + 2) / 5 + 365 * y + y / 4 - y / 100 + y / 400 - 32045;
let days_since_epoch = jdn - 2_440_588;
Some(days_since_epoch * 86_400 + hour * 3_600 + min * 60 + sec)
}
fn spawn_connection_task<R, W>(
reader: R,
writer: W,
params: ConnectionLoopParams,
alive: Arc<std::sync::atomic::AtomicBool>,
) -> tokio::task::JoinHandle<()>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
let close_metrics = params.metrics.clone();
tokio::spawn(async move {
let result =
std::panic::AssertUnwindSafe(Self::run_connection_loop(reader, writer, params))
.catch_unwind()
.await;
match result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
close_metrics.record_error();
error!("Connection error: {e}");
}
Err(_panic_payload) => {
close_metrics.record_error();
error!("Connection event loop panicked; all in-flight requests failed");
}
}
close_metrics.record_close();
alive.store(false, std::sync::atomic::Ordering::Release);
})
}
async fn run_connection_loop<R, W>(
mut reader: R,
mut writer: W,
params: ConnectionLoopParams,
) -> Result<()>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
let ConnectionLoopParams {
address: broker_address,
mut high_priority_rx,
mut normal_priority_rx,
request_timeout,
stats,
metrics,
max_response_size,
max_in_flight_requests,
max_high_priority_bypasses: max_high_priority_bypasses_per_round,
} = params;
let mut pending: AHashMap<CorrelationId, PendingRequest> = AHashMap::new();
let mut delay_queue: DelayQueue<CorrelationId> = DelayQueue::new();
let mut delay_keys: AHashMap<CorrelationId, delay_queue::Key> = AHashMap::new();
let mut timed_out: std::collections::VecDeque<CorrelationId> =
std::collections::VecDeque::new();
let (frame_tx, mut frame_rx) =
mpsc::channel::<Result<Bytes>>(max_in_flight_requests.max(1));
let reader_handle = tokio::spawn(async move {
let mut decoder = Decoder::with_max_size(max_response_size);
let mut buf = vec![0u8; 65536];
loop {
match reader.read(&mut buf).await {
Ok(0) => {
debug!("Connection closed by peer");
break;
}
Ok(n) => {
decoder.extend(&buf[..n]);
loop {
match decoder.decode() {
Ok(Some(frame)) => {
if frame_tx.send(Ok(frame)).await.is_err() {
return Ok::<_, KrafkaError>(());
}
}
Ok(None) => break,
Err(e) => {
let _ = frame_tx.send(Err(e)).await;
return Ok(());
}
}
}
}
Err(e) => {
let _ = frame_tx.send(Err(KrafkaError::network(e))).await;
return Ok(());
}
}
}
Ok(())
});
let mut terminal_error: Option<KrafkaError> = None;
let mut consecutive_high_priority_commands = 0usize;
let mut deferred_high_priority_cmd: Option<ConnectionCommand> = None;
loop {
if consecutive_high_priority_commands >= max_high_priority_bypasses_per_round {
if deferred_high_priority_cmd.is_none() {
match high_priority_rx.try_recv() {
Ok(ConnectionCommand::Close) => {
consecutive_high_priority_commands = 0;
match Self::process_loop_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
ConnectionCommand::Close,
max_in_flight_requests,
request_timeout,
)
.await
{
Ok(true) => break,
Ok(false) => {}
Err(err) => {
terminal_error = Some(err);
break;
}
}
continue;
}
Ok(cmd) => {
deferred_high_priority_cmd = Some(cmd);
}
Err(mpsc::error::TryRecvError::Empty)
| Err(mpsc::error::TryRecvError::Disconnected) => {}
}
}
match normal_priority_rx.try_recv() {
Ok(cmd) => {
stats
.high_priority_bypass_yields
.fetch_add(1, Ordering::Relaxed);
metrics.record_high_priority_bypass_yield();
consecutive_high_priority_commands = 0;
match Self::process_loop_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await
{
Ok(true) => break,
Ok(false) => {}
Err(err) => {
terminal_error = Some(err);
break;
}
}
continue;
}
Err(mpsc::error::TryRecvError::Empty)
| Err(mpsc::error::TryRecvError::Disconnected) => {
consecutive_high_priority_commands = 0;
}
}
}
if let Some(cmd) = deferred_high_priority_cmd.take() {
consecutive_high_priority_commands += 1;
match Self::process_loop_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await
{
Ok(true) => break,
Ok(false) => {}
Err(err) => {
terminal_error = Some(err);
break;
}
}
continue;
}
if let Ok(cmd) = high_priority_rx.try_recv() {
stats.high_priority_bypasses.fetch_add(1, Ordering::Relaxed);
metrics.record_high_priority_bypass();
consecutive_high_priority_commands += 1;
match Self::process_loop_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await
{
Ok(true) => break,
Ok(false) => {}
Err(err) => {
terminal_error = Some(err);
break;
}
}
continue;
}
tokio::select! {
biased;
frame_result = frame_rx.recv() => {
consecutive_high_priority_commands = 0;
match frame_result {
Some(Ok(frame)) => {
if let Err(e) = Self::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&mut timed_out,
frame,
&broker_address,
) {
terminal_error = Some(e);
break;
}
}
Some(Err(e)) => {
terminal_error = Some(e);
break;
}
None => {
break;
}
}
}
Some(expired) = std::future::poll_fn(|cx| {
use futures_core::Stream;
std::pin::Pin::new(&mut delay_queue).poll_next(cx)
}) => {
consecutive_high_priority_commands = 0;
let id = expired.into_inner();
if let Some(req) = pending.remove(&id) {
delay_keys.remove(&id);
if timed_out.len() >= max_in_flight_requests {
timed_out.pop_front();
}
timed_out.push_back(id);
warn!(
correlation_id = id,
"Request timed out after {:?}", request_timeout
);
let _ = req.response_tx.send(Err(KrafkaError::timeout(format!(
"request {id} timed out after {request_timeout:?}"
))));
}
}
cmd = high_priority_rx.recv() => {
match cmd {
Some(cmd) => {
consecutive_high_priority_commands += 1;
match Self::process_loop_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await {
Ok(true) => break,
Ok(false) => {}
Err(err) => {
terminal_error = Some(err);
break;
}
}
}
None => break,
}
}
cmd = normal_priority_rx.recv() => {
match cmd {
Some(cmd) => {
consecutive_high_priority_commands = 0;
match Self::process_loop_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await {
Ok(true) => break,
Ok(false) => {}
Err(err) => {
terminal_error = Some(err);
break;
}
}
}
None => break,
}
}
}
}
let _ = writer.shutdown().await;
drop(writer);
reader_handle.abort();
let pending_error = terminal_error
.clone()
.unwrap_or_else(connection_closed_error);
for (_, req) in pending.drain() {
let _ = req.response_tx.send(Err(pending_error.clone()));
}
if let Some(err) = terminal_error {
return Err(err);
}
Ok(())
}
async fn process_loop_command<W: AsyncWrite + Unpin>(
writer: &mut W,
pending: &mut AHashMap<CorrelationId, PendingRequest>,
delay_queue: &mut DelayQueue<CorrelationId>,
delay_keys: &mut AHashMap<CorrelationId, delay_queue::Key>,
cmd: ConnectionCommand,
max_in_flight_requests: usize,
request_timeout: Duration,
) -> Result<bool> {
Self::handle_command_direct(
writer,
pending,
delay_queue,
delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await
}
async fn handle_command_direct<W: AsyncWrite + Unpin>(
writer: &mut W,
pending: &mut AHashMap<CorrelationId, PendingRequest>,
delay_queue: &mut DelayQueue<CorrelationId>,
delay_keys: &mut AHashMap<CorrelationId, delay_queue::Key>,
cmd: ConnectionCommand,
max_in_flight_requests: usize,
request_timeout: Duration,
) -> Result<bool> {
match cmd {
ConnectionCommand::Request {
data,
correlation_id,
api_key,
api_version,
response_tx,
timeout: budget,
permit,
} => {
if pending.contains_key(&correlation_id) {
let error = KrafkaError::invalid_state(format!(
"correlation ID collision on broker connection: correlation_id={correlation_id}, pending_requests={}; closing connection",
pending.len()
));
error!(
correlation_id,
pending_requests = pending.len(),
"Detected correlation ID collision; closing connection"
);
let _ = response_tx.send(Err(error.clone()));
return Err(error);
}
if pending.len() >= max_in_flight_requests {
warn!(
pending = pending.len(),
max = max_in_flight_requests,
"Rejecting request: max in-flight requests reached \
(in-flight permit accounting inconsistent)"
);
let _ = response_tx.send(Err(KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::WouldBlock,
format!("max in-flight requests ({max_in_flight_requests}) reached; retry"),
))));
return Ok(false);
}
let deadline = tokio::time::Instant::now() + budget;
let write_result = tokio::time::timeout_at(deadline, async {
writer.write_all(&data).await?;
writer.flush().await
})
.await;
match write_result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
error!("Write error: {}", e);
let msg = e.to_string();
let _ = response_tx.send(Err(KrafkaError::network(e)));
return Err(KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
format!("write failed, stream indeterminate: {msg}"),
)));
}
Err(_) => {
let msg = format!("write timed out after {budget:?}");
error!("{msg}");
let _ = response_tx.send(Err(KrafkaError::timeout(msg.clone())));
return Err(KrafkaError::timeout(msg));
}
}
let key = delay_queue.insert_at(correlation_id, deadline);
delay_keys.insert(correlation_id, key);
pending.insert(
correlation_id,
PendingRequest {
response_tx,
api_key,
api_version,
_permit: permit,
},
);
Ok(false)
}
ConnectionCommand::Close => {
debug!("Closing connection");
Ok(true)
}
ConnectionCommand::FireAndForget { data } => {
let write_result = tokio::time::timeout(request_timeout, async {
writer.write_all(&data).await?;
writer.flush().await
})
.await;
match write_result {
Ok(Ok(())) => {}
Ok(Err(e)) => {
error!("Fire-and-forget write error: {}", e);
return Err(KrafkaError::network(e));
}
Err(_) => {
error!(
"Fire-and-forget write timed out after {:?}",
request_timeout
);
return Err(KrafkaError::timeout(format!(
"fire-and-forget write timed out after {request_timeout:?}"
)));
}
}
Ok(false)
}
}
}
fn dispatch_response(
pending: &mut AHashMap<CorrelationId, PendingRequest>,
delay_queue: &mut DelayQueue<CorrelationId>,
delay_keys: &mut AHashMap<CorrelationId, delay_queue::Key>,
timed_out: &mut std::collections::VecDeque<CorrelationId>,
response: Bytes,
broker_address: &str,
) -> Result<()> {
if response.len() < 4 {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::TruncatedFrame,
format!(
"response too short from broker {broker_address}: frame_bytes={}",
response.len()
),
));
}
let correlation_id =
i32::from_be_bytes([response[0], response[1], response[2], response[3]]);
let pending_before_remove = pending.len();
if let Some(req) = pending.remove(&correlation_id) {
if let Some(key) = delay_keys.remove(&correlation_id) {
delay_queue.remove(&key);
}
trace!("Received response for correlation_id={}", correlation_id);
let mut response_buf = response.slice(..);
match ResponseHeader::decode(&mut response_buf, req.api_key, req.api_version) {
Ok(_header) => {
let header_size = response.len() - response_buf.len();
let body = response.slice(header_size..);
let _ = req.response_tx.send(Ok(body));
}
Err(e) => {
let response_header_version =
ResponseHeader::header_version(req.api_key, req.api_version);
let context = format!(
"response header decode failed: broker={broker_address}, api_key={:?}, api_version={}, response_header_version={}, correlation_id={correlation_id}, frame_bytes={}, pending_before_remove={pending_before_remove}, error={e}",
req.api_key,
req.api_version,
response_header_version,
response.len(),
);
warn!(
broker = broker_address,
api_key = ?req.api_key,
api_version = req.api_version,
response_header_version,
correlation_id,
frame_bytes = response.len(),
pending_before_remove,
error = %e,
"Failed to decode response header; closing connection"
);
let _ = req.response_tx.send(Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
context.clone(),
)));
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!("{context}; stream desynchronized"),
));
}
}
} else if let Some(pos) = timed_out.iter().position(|&id| id == correlation_id) {
timed_out.remove(pos);
debug!(
correlation_id,
broker = broker_address,
frame_bytes = response.len(),
"Discarding late response for a timed-out request"
);
} else {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!(
"Received response for unknown correlation_id={correlation_id} from broker {broker_address}; frame_bytes={}, pending_requests={pending_before_remove}; closing connection",
response.len()
),
));
}
Ok(())
}
async fn acquire_in_flight(
&self,
deadline: tokio::time::Instant,
) -> Result<tokio::sync::OwnedSemaphorePermit> {
tokio::time::timeout_at(deadline, self.in_flight.clone().acquire_owned())
.await
.map_err(|_| {
KrafkaError::timeout(format!(
"waiting for an in-flight slot on {} (max_in_flight_requests={})",
self.address, self.config.max_in_flight_requests
))
})?
.map_err(|_| connection_closed_error())
}
async fn fetch_api_versions(&self) -> Result<()> {
const MAX_ATTEMPTS: usize = 3;
let request =
ApiVersionsRequest::new().with_client_software("krafka", env!("CARGO_PKG_VERSION"));
let mut attempt_version = crate::protocol::versions::API_VERSIONS_MAX;
let mut attempts = 0usize;
let response = loop {
attempts += 1;
let body = self.send_api_versions(&request, attempt_version).await?;
let error_code = body
.get(..2)
.map(|b| i16::from_be_bytes([b[0], b[1]]))
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::TruncatedFrame,
"ApiVersions response too short to contain an error code",
)
})?;
let unsupported = error_code == ErrorCode::UnsupportedVersion.to_i16();
let mut buf = body;
let decoded = if unsupported {
ApiVersionsResponse::decode_v0(&mut buf)?
} else {
Self::decode_api_versions_at(attempt_version, &mut buf)?
};
if !unsupported {
break decoded;
}
let next = decoded
.get_api_version(ApiKey::ApiVersions)
.map(|range| range.max_version)
.filter(|&max| max >= 0 && max < attempt_version)
.unwrap_or(0);
if next >= attempt_version || attempts >= MAX_ATTEMPTS {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
format!(
"broker {} rejected ApiVersions v{attempt_version} with \
UNSUPPORTED_VERSION and offered no lower usable version",
self.address
),
));
}
debug!(
broker = %self.address,
rejected = attempt_version,
retrying_with = next,
"Broker rejected the ApiVersions version; falling back"
);
attempt_version = next;
};
if response.error_code != 0 {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Other,
format!("ApiVersions error: {}", response.error_code),
));
}
{
let mut features = self.broker_features.lock();
features.negotiated_version = attempt_version;
features.supported.clone_from(&response.supported_features);
features.finalized_epoch = response.finalized_features_epoch;
features.finalized.clone_from(&response.finalized_features);
}
let mut versions = self.api_versions.lock();
for range in response.api_keys {
versions.insert(range.api_key, range);
}
debug!(
broker = %self.address,
api_versions_version = attempt_version,
apis = versions.len(),
"Negotiated broker API versions"
);
Ok(())
}
fn decode_api_versions_at(version: i16, buf: &mut Bytes) -> Result<ApiVersionsResponse> {
match version {
0 => ApiVersionsResponse::decode_v0(buf),
1..=2 => ApiVersionsResponse::decode_v1(buf),
_ => ApiVersionsResponse::decode_v3(buf),
}
}
async fn send_api_versions(&self, request: &ApiVersionsRequest, version: i16) -> Result<Bytes> {
let correlation_id = self.correlation_id_gen.next();
let mut encoder = Encoder::with_capacity(128);
let pos = encoder.start_message();
let header = RequestHeader::new(ApiKey::ApiVersions, version, correlation_id)
.with_client_id(&self.config.client_id);
header.encode(encoder.buffer_mut())?;
match version {
0..=2 => request.encode_v0(encoder.buffer_mut())?,
3..=4 => request.encode_v3(encoder.buffer_mut())?,
_ => request.encode_v5(encoder.buffer_mut())?,
}
encoder.finish_message(pos)?;
let deadline = tokio::time::Instant::now() + self.config.request_timeout;
let permit = self.acquire_in_flight(deadline).await?;
let (response_tx, response_rx) = oneshot::channel();
tokio::time::timeout_at(
deadline,
self.high_priority_tx.send(ConnectionCommand::Request {
data: encoder.take(),
correlation_id,
api_key: ApiKey::ApiVersions,
api_version: version,
response_tx,
timeout: self.config.request_timeout,
permit,
}),
)
.await
.map_err(|_| KrafkaError::timeout("enqueuing api versions request"))?
.map_err(|_| connection_closed_error())?;
self.stats
.high_priority_requests
.fetch_add(1, Ordering::Relaxed);
self.config
.connection_metrics
.record_high_priority_request();
tokio::time::timeout_at(deadline, response_rx)
.await
.map_err(|_| KrafkaError::timeout("api versions request"))?
.map_err(|_| connection_closed_error())?
}
#[must_use]
pub fn broker_features(&self) -> BrokerFeatures {
self.broker_features.lock().clone()
}
#[inline]
fn channel_for_priority(&self, priority: RequestPriority) -> &mpsc::Sender<ConnectionCommand> {
match priority {
RequestPriority::High => &self.high_priority_tx,
RequestPriority::Normal => &self.normal_priority_tx,
}
}
pub fn notify_throttle(&self, throttle_time_ms: i32) {
if throttle_time_ms > 0 {
let new_deadline = Instant::now() + Duration::from_millis(throttle_time_ms as u64);
let mut deadline = self.throttle_until.lock();
if new_deadline > *deadline {
debug!(
throttle_ms = throttle_time_ms,
broker = %self.address,
"Broker throttle applied (KIP-219)"
);
*deadline = new_deadline;
}
}
}
#[inline]
pub fn throttle_remaining(&self) -> Option<Duration> {
self.throttle_until
.lock()
.checked_duration_since(Instant::now())
}
pub async fn await_throttle(&self) -> Option<Duration> {
let remaining = self.throttle_remaining()?;
debug!(
delay_ms = remaining.as_millis() as u64,
broker = %self.address,
"Delaying request due to broker throttle (KIP-219)"
);
self.config
.connection_metrics
.record_throttle_delay(remaining);
tokio::time::sleep(remaining).await;
Some(remaining)
}
pub async fn send_request(
&self,
api_key: ApiKey,
api_version: i16,
request_body: impl FnOnce(&mut BytesMut) -> Result<()>,
) -> Result<Bytes> {
let priority = RequestPriority::for_api_key(api_key);
self.send_request_with_priority(api_key, api_version, priority, request_body)
.await
}
pub async fn send_request_with_timeout(
&self,
api_key: ApiKey,
api_version: i16,
timeout: Duration,
request_body: impl FnOnce(&mut BytesMut) -> Result<()>,
) -> Result<Bytes> {
let priority = RequestPriority::for_api_key(api_key);
let budget = timeout.max(self.config.request_timeout);
self.send_inner(api_key, api_version, priority, budget, request_body)
.await
}
pub async fn send_request_with_priority(
&self,
api_key: ApiKey,
api_version: i16,
priority: RequestPriority,
request_body: impl FnOnce(&mut BytesMut) -> Result<()>,
) -> Result<Bytes> {
self.send_inner(
api_key,
api_version,
priority,
self.config.request_timeout,
request_body,
)
.await
}
async fn send_inner(
&self,
api_key: ApiKey,
api_version: i16,
priority: RequestPriority,
budget: Duration,
request_body: impl FnOnce(&mut BytesMut) -> Result<()>,
) -> Result<Bytes> {
self.mark_used();
let deadline = tokio::time::Instant::now() + budget;
if priority == RequestPriority::Normal {
self.await_throttle().await;
}
let correlation_id = self.correlation_id_gen.next();
let mut encoder = Encoder::with_capacity(256);
let pos = encoder.start_message();
let header = RequestHeader::new(api_key, api_version, correlation_id)
.with_client_id(&self.config.client_id);
header.encode(encoder.buffer_mut())?;
request_body(encoder.buffer_mut())?;
encoder.finish_message(pos)?;
let permit = self.acquire_in_flight(deadline).await?;
let (response_tx, response_rx) = oneshot::channel();
let channel = self.channel_for_priority(priority);
tokio::time::timeout_at(
deadline,
channel.send(ConnectionCommand::Request {
data: encoder.take(),
correlation_id,
api_key,
api_version,
response_tx,
timeout: budget,
permit,
}),
)
.await
.map_err(|_| {
KrafkaError::timeout(format!(
"enqueuing {api_key:?} request to {} (channel full)",
self.address
))
})?
.map_err(|_| connection_closed_error())?;
match priority {
RequestPriority::High => {
self.stats
.high_priority_requests
.fetch_add(1, Ordering::Relaxed);
self.config
.connection_metrics
.record_high_priority_request();
}
RequestPriority::Normal => {
self.stats
.normal_priority_requests
.fetch_add(1, Ordering::Relaxed);
self.config
.connection_metrics
.record_normal_priority_request();
}
}
let response = tokio::time::timeout_at(deadline, response_rx)
.await
.map_err(|_| KrafkaError::timeout("request"))?
.map_err(|_| connection_closed_error())??;
if let Some(throttle_time_ms) = leading_throttle_time_ms(api_key, api_version, &response) {
self.notify_throttle(throttle_time_ms);
}
Ok(response)
}
pub async fn send_fire_and_forget(
&self,
api_key: ApiKey,
api_version: i16,
request_body: impl FnOnce(&mut BytesMut) -> Result<()>,
) -> Result<()> {
self.mark_used();
let mut encoder = Encoder::with_capacity(256);
let pos = encoder.start_message();
let header = RequestHeader::new(api_key, api_version, NO_RESPONSE_CORRELATION_ID)
.with_client_id(&self.config.client_id);
header.encode(encoder.buffer_mut())?;
request_body(encoder.buffer_mut())?;
encoder.finish_message(pos)?;
let channel = self.channel_for_priority(RequestPriority::Normal);
tokio::time::timeout(
self.config.request_timeout,
channel.send(ConnectionCommand::FireAndForget {
data: encoder.take(),
}),
)
.await
.map_err(|_| {
KrafkaError::timeout(format!(
"enqueuing fire-and-forget {api_key:?} to {} (channel full)",
self.address
))
})?
.map_err(|_| connection_closed_error())?;
self.stats
.normal_priority_requests
.fetch_add(1, Ordering::Relaxed);
self.config
.connection_metrics
.record_normal_priority_request();
Ok(())
}
pub fn get_api_version(&self, api_key: ApiKey) -> Option<ApiVersionRange> {
let versions = self.api_versions.lock();
versions.get(&api_key).copied()
}
pub fn negotiate_api_version(
&self,
api_key: ApiKey,
client_max: i16,
client_min: i16,
) -> Option<i16> {
let versions = self.api_versions.lock();
versions
.get(&api_key)
.and_then(|range| range.negotiate(client_max, client_min))
}
pub fn negotiate_api_version_max(&self, api_key: ApiKey, client_max: i16) -> Option<i16> {
self.negotiate_api_version(api_key, client_max, 0)
}
fn compute_session_expiry(session_lifetime_ms: i64) -> Option<Instant> {
if session_lifetime_ms <= 0 {
return None;
}
const MIN_REAUTH_MS: u64 = 100;
let base_factor: f64 = 0.85;
let jitter_range: f64 = 0.10;
let jitter: f64 = rand::random::<f64>() * jitter_range;
let factor = base_factor + jitter;
let computed_reauth_ms = (session_lifetime_ms as f64 * factor) as u64;
let reauth_ms = computed_reauth_ms.max(MIN_REAUTH_MS);
if computed_reauth_ms < MIN_REAUTH_MS {
warn!(
session_lifetime_ms,
computed_reauth_ms,
reauth_ms,
"broker reported unusually small SASL session lifetime; clamping reauthentication delay"
);
}
Some(Instant::now() + Duration::from_millis(reauth_ms))
}
fn effective_session_expiry(session_lifetime_ms: i64, auth: &AuthConfig) -> Option<Instant> {
if session_lifetime_ms > 0 {
return Self::compute_session_expiry(session_lifetime_ms);
}
if let Some(token) = auth.oauthbearer_token.as_ref()
&& let Some(expiry_epoch_ms) = token.lifetime_ms()
{
let now_epoch_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as i64;
let remaining_ms = expiry_epoch_ms.saturating_sub(now_epoch_ms);
if remaining_ms > 0 {
return Self::compute_session_expiry(remaining_ms);
}
}
None
}
#[inline]
pub fn needs_reauthentication(&self) -> bool {
self.session_expiry
.is_some_and(|expiry| Instant::now() >= expiry)
}
#[inline]
pub fn session_expiry(&self) -> Option<Instant> {
self.session_expiry
}
#[inline]
pub fn is_alive(&self) -> bool {
self.alive.load(std::sync::atomic::Ordering::Acquire)
}
#[inline]
pub fn is_usable(&self) -> bool {
self.is_alive() && !self.needs_reauthentication()
}
#[inline]
fn mark_used(&self) {
let elapsed = self.created_at.elapsed().as_nanos();
let nanos = u64::try_from(elapsed).unwrap_or(u64::MAX);
self.last_used_nanos.store(nanos, Ordering::Relaxed);
}
#[inline]
pub fn idle_duration(&self) -> Duration {
let last = self.last_used_nanos.load(Ordering::Relaxed);
let now = self.created_at.elapsed();
now.saturating_sub(Duration::from_nanos(last))
}
#[cfg(test)]
#[allow(clippy::expect_used)]
pub(crate) fn test_stub_idle_for(address: &str, idle_for: Duration) -> Self {
let (high_priority_tx, _) = mpsc::channel(1);
let (normal_priority_tx, _) = mpsc::channel(1);
Self {
address: address.to_string(),
config: ConnectionConfig::default(),
correlation_id_gen: Arc::new(CorrelationIdGenerator::new()),
high_priority_tx,
normal_priority_tx,
api_versions: Arc::new(parking_lot::Mutex::new(AHashMap::new())),
broker_features: Arc::new(parking_lot::Mutex::new(BrokerFeatures::default())),
alive: Arc::new(std::sync::atomic::AtomicBool::new(true)),
session_expiry: None,
stats: Arc::new(ConnectionStats::default()),
throttle_until: Arc::new(parking_lot::Mutex::new(Instant::now())),
created_at: Instant::now()
.checked_sub(idle_for)
.expect("idle_for exceeds system uptime; cannot backdate Instant"),
last_used_nanos: AtomicU64::new(0),
in_flight: Arc::new(tokio::sync::Semaphore::new(1)),
}
}
#[cfg(test)]
pub(crate) fn test_mark_fresh(&self) {
self.mark_used();
}
#[inline]
pub fn address(&self) -> &str {
&self.address
}
#[inline]
pub fn stats(&self) -> &ConnectionStats {
&self.stats
}
pub async fn close(&self) {
match self.high_priority_tx.try_send(ConnectionCommand::Close) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Closed(_)) => {}
Err(mpsc::error::TrySendError::Full(cmd)) => {
if timeout(CLOSE_SEND_TIMEOUT, self.high_priority_tx.send(cmd))
.await
.is_err()
{
warn!(
broker = %self.address,
"close() timed out after {CLOSE_SEND_TIMEOUT:?} waiting for the \
high-priority channel; the event loop is stalled. Socket teardown \
falls back to Drop."
);
}
}
}
}
}
impl Drop for BrokerConnection {
fn drop(&mut self) {
if let Ok(_handle) = tokio::runtime::Handle::try_current() {
let tx = self.high_priority_tx.clone();
tokio::spawn(async move {
let _ = tx.send(ConnectionCommand::Close).await;
});
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
fn mock_api_versions_body(api_version: i16) -> BytesMut {
use bytes::BufMut as _;
let mut body = BytesMut::new();
body.put_i16(0); if api_version >= 3 {
body.put_u8(1);
body.put_i32(0); body.put_u8(0); } else {
body.put_i32(0); if api_version >= 1 {
body.put_i32(0); }
}
body
}
fn test_permit() -> tokio::sync::OwnedSemaphorePermit {
Arc::new(tokio::sync::Semaphore::new(1))
.try_acquire_owned()
.expect("fresh semaphore always has a permit")
}
#[test]
fn test_connection_config_builder() {
let config = ConnectionConfig::builder()
.connect_timeout(Duration::from_secs(5))
.request_timeout(Duration::from_secs(15))
.client_id("test-client")
.nodelay(false)
.build()
.unwrap();
assert_eq!(config.connect_timeout, Duration::from_secs(5));
assert_eq!(config.request_timeout, Duration::from_secs(15));
assert_eq!(config.client_id, "test-client");
assert!(!config.nodelay);
}
#[test]
fn test_connection_config_default() {
let config = ConnectionConfig::default();
assert_eq!(config.connect_timeout, Duration::from_secs(10));
assert_eq!(config.request_timeout, Duration::from_secs(30));
assert_eq!(config.client_id, "krafka");
assert!(config.nodelay);
assert_eq!(config.high_priority_channel_capacity, 64);
assert_eq!(config.normal_priority_channel_capacity, 256);
assert!(config.auth.is_none());
}
#[test]
fn test_connection_config_uses_shared_metrics_handle() {
let metrics = Arc::new(ConnectionMetrics::default());
let config = ConnectionConfig::builder()
.connection_metrics(metrics.clone())
.build()
.unwrap();
config.connection_metrics.record_high_priority_request();
assert_eq!(metrics.high_priority_requests.get(), 1);
assert!(Arc::ptr_eq(&metrics, &config.connection_metrics()));
}
#[test]
fn test_connection_config_with_auth() {
use crate::auth::AuthConfig;
let config = ConnectionConfig::builder()
.client_id("test")
.auth(AuthConfig::sasl_plain("user", "pass").unwrap())
.build()
.unwrap();
assert_eq!(config.client_id, "test");
let auth = config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(!auth.requires_tls());
}
#[test]
fn test_connection_config_builder_with_priority() {
let config = ConnectionConfig::builder()
.high_priority_channel_capacity(32)
.normal_priority_channel_capacity(512)
.build()
.unwrap();
assert_eq!(config.high_priority_channel_capacity, 32);
assert_eq!(config.normal_priority_channel_capacity, 512);
}
#[test]
fn test_connection_config_min_values() {
let config = ConnectionConfig::builder()
.high_priority_channel_capacity(0) .normal_priority_channel_capacity(0) .build()
.unwrap();
assert_eq!(config.high_priority_channel_capacity, 16);
assert_eq!(config.normal_priority_channel_capacity, 64);
}
#[test]
fn test_request_priority_for_api_key() {
assert_eq!(
RequestPriority::for_api_key(ApiKey::Heartbeat),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::Metadata),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::FindCoordinator),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::ApiVersions),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::ConsumerGroupHeartbeat),
RequestPriority::High,
"ConsumerGroupHeartbeat must be High to prevent KIP-848 rebalances"
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::ShareGroupHeartbeat),
RequestPriority::High,
"ShareGroupHeartbeat must be High to prevent KIP-932 share group evictions"
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::JoinGroup),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::SyncGroup),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::LeaveGroup),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::OffsetCommit),
RequestPriority::High
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::Produce),
RequestPriority::Normal
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::Fetch),
RequestPriority::Normal
);
assert_eq!(
RequestPriority::for_api_key(ApiKey::OffsetFetch),
RequestPriority::Normal
);
}
#[test]
fn test_connection_stats_default() {
let stats = ConnectionStats::default();
assert_eq!(stats.high_priority_count(), 0);
assert_eq!(stats.normal_priority_count(), 0);
assert_eq!(stats.bypass_count(), 0);
assert_eq!(stats.bypass_yield_count(), 0);
}
#[test]
fn test_connection_stats_increment() {
let stats = ConnectionStats::default();
stats.high_priority_requests.fetch_add(5, Ordering::Relaxed);
stats
.normal_priority_requests
.fetch_add(10, Ordering::Relaxed);
stats.high_priority_bypasses.fetch_add(2, Ordering::Relaxed);
stats
.high_priority_bypass_yields
.fetch_add(1, Ordering::Relaxed);
assert_eq!(stats.high_priority_count(), 5);
assert_eq!(stats.normal_priority_count(), 10);
assert_eq!(stats.bypass_count(), 2);
assert_eq!(stats.bypass_yield_count(), 1);
}
#[test]
fn test_dispatch_response_late_response_does_not_desync_connection() {
let correlation_id: CorrelationId = 42;
let mut pending = AHashMap::new(); let mut delay_queue = DelayQueue::new();
let mut delay_keys = AHashMap::new();
let mut timed_out = std::collections::VecDeque::from([correlation_id]);
let result = BrokerConnection::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&mut timed_out,
Bytes::copy_from_slice(&correlation_id.to_be_bytes()),
"broker-1:9092",
);
assert!(
result.is_ok(),
"late response must be discarded, not reported as desync: {result:?}"
);
assert!(timed_out.is_empty());
let second = BrokerConnection::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&mut timed_out,
Bytes::copy_from_slice(&correlation_id.to_be_bytes()),
"broker-1:9092",
);
assert!(
second.is_err(),
"a truly unknown correlation id is still fatal"
);
}
#[test]
fn test_dispatch_response_header_decode_error_includes_context() {
let correlation_id = 7;
let (response_tx, mut response_rx) = oneshot::channel();
let mut pending = AHashMap::new();
pending.insert(
correlation_id,
PendingRequest {
response_tx,
api_key: ApiKey::Metadata,
api_version: 9,
_permit: test_permit(),
},
);
let mut delay_queue = DelayQueue::new();
let mut delay_keys = AHashMap::new();
let mut timed_out = std::collections::VecDeque::new();
let err = BrokerConnection::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&mut timed_out,
Bytes::copy_from_slice(&correlation_id.to_be_bytes()),
"broker-1:9092",
)
.unwrap_err();
let caller_err = response_rx.try_recv().unwrap().unwrap_err();
let err_text = caller_err.to_string();
assert!(err.to_string().contains("stream desynchronized"));
assert!(err_text.contains("broker=broker-1:9092"));
assert!(err_text.contains("api_key=Metadata"));
assert!(err_text.contains("api_version=9"));
assert!(err_text.contains("response_header_version=1"));
assert!(err_text.contains("correlation_id=7"));
assert!(err_text.contains("frame_bytes=4"));
}
async fn run_mock_sasl_broker(
listener: tokio::net::TcpListener,
shutdown_rx: oneshot::Receiver<()>,
) -> (String, Vec<u8>) {
run_mock_sasl_broker_with_lifetime(listener, 0, shutdown_rx).await
}
async fn run_mock_sasl_broker_with_lifetime(
listener: tokio::net::TcpListener,
session_lifetime_ms: i64,
shutdown_rx: oneshot::Receiver<()>,
) -> (String, Vec<u8>) {
use bytes::BufMut;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut stream, _) = listener.accept().await.unwrap();
async fn read_frame(stream: &mut tokio::net::TcpStream) -> Vec<u8> {
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await.unwrap();
let len = i32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
stream.read_exact(&mut body).await.unwrap();
body
}
async fn write_frame(stream: &mut tokio::net::TcpStream, data: &[u8]) {
let len = data.len() as i32;
stream.write_all(&len.to_be_bytes()).await.unwrap();
stream.write_all(data).await.unwrap();
stream.flush().await.unwrap();
}
let req = read_frame(&mut stream).await;
let correlation_id = i32::from_be_bytes(req[4..8].try_into().unwrap());
let client_id_len = i16::from_be_bytes(req[8..10].try_into().unwrap());
let mech_offset = if client_id_len < 0 {
10 } else {
10 + client_id_len as usize
};
let mech_len =
i16::from_be_bytes(req[mech_offset..mech_offset + 2].try_into().unwrap()) as usize;
let mechanism =
String::from_utf8(req[mech_offset + 2..mech_offset + 2 + mech_len].to_vec()).unwrap();
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_i16(0); resp.put_i32(1); let mech_bytes = mechanism.as_bytes();
resp.put_i16(mech_bytes.len() as i16);
resp.put_slice(mech_bytes);
write_frame(&mut stream, &resp).await;
let req = read_frame(&mut stream).await;
let correlation_id = i32::from_be_bytes(req[4..8].try_into().unwrap());
let client_id_len = i16::from_be_bytes(req[8..10].try_into().unwrap());
let auth_offset = if client_id_len < 0 {
10
} else {
10 + client_id_len as usize
};
let auth_bytes_len =
i32::from_be_bytes(req[auth_offset..auth_offset + 4].try_into().unwrap()) as usize;
let auth_bytes = req[auth_offset + 4..auth_offset + 4 + auth_bytes_len].to_vec();
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_i16(0); resp.put_i16(-1_i16); resp.put_i32(0); resp.put_i64(session_lifetime_ms); write_frame(&mut stream, &resp).await;
let req = read_frame(&mut stream).await;
let api_version = i16::from_be_bytes(req[2..4].try_into().unwrap());
let correlation_id = i32::from_be_bytes(req[4..8].try_into().unwrap());
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_slice(&mock_api_versions_body(api_version));
write_frame(&mut stream, &resp).await;
let _ = shutdown_rx.await;
(mechanism, auth_bytes)
}
#[tokio::test]
async fn test_sasl_plain_handshake_with_mock_broker() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let addr_str = addr.to_string();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let mock_handle = tokio::spawn(run_mock_sasl_broker(listener, shutdown_rx));
let config = ConnectionConfig::builder()
.client_id("test-client")
.auth(crate::auth::AuthConfig::sasl_plain("testuser", "testpassword").unwrap())
.build()
.unwrap();
let conn = BrokerConnection::connect(&addr_str, config).await;
assert!(
conn.is_ok(),
"Connection with SASL/PLAIN should succeed: {:?}",
conn.err()
);
let conn = conn.unwrap();
assert!(conn.is_alive());
conn.close().await;
let _ = shutdown_tx.send(());
let (mechanism, auth_bytes) = mock_handle.await.unwrap();
assert_eq!(mechanism, "PLAIN");
assert_eq!(auth_bytes, b"\0testuser\0testpassword");
}
#[tokio::test]
async fn test_sasl_oauthbearer_provider_handshake_with_mock_broker() {
use crate::auth::OAuthBearerToken;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let addr_str = addr.to_string();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let mock_handle = tokio::spawn(run_mock_sasl_broker(listener, shutdown_rx));
let config = ConnectionConfig::builder()
.client_id("test-client")
.auth(crate::auth::AuthConfig::sasl_oauthbearer_provider(
|| async { Ok(OAuthBearerToken::new("provider-jwt-token")) },
))
.build()
.unwrap();
let conn = BrokerConnection::connect(&addr_str, config).await;
assert!(
conn.is_ok(),
"Connection with OAUTHBEARER provider should succeed: {:?}",
conn.err()
);
let conn = conn.unwrap();
assert!(conn.is_alive());
conn.close().await;
let _ = shutdown_tx.send(());
let (mechanism, auth_bytes) = mock_handle.await.unwrap();
assert_eq!(mechanism, "OAUTHBEARER");
let expected = OAuthBearerToken::new("provider-jwt-token").to_gs2_initial_response();
assert_eq!(auth_bytes, expected);
}
#[tokio::test]
async fn test_sasl_oauthbearer_provider_timeout() {
let config = ConnectionConfig::builder()
.client_id("test-client")
.connect_timeout(Duration::from_millis(50))
.request_timeout(Duration::from_millis(100))
.auth(crate::auth::AuthConfig::sasl_oauthbearer_provider(
|| async {
tokio::time::sleep(Duration::from_secs(60)).await;
Ok(crate::auth::OAuthBearerToken::new("never"))
},
))
.build()
.unwrap();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let addr_str = addr.to_string();
tokio::spawn(async move {
let (_stream, _) = listener.accept().await.unwrap();
tokio::time::sleep(Duration::from_secs(5)).await;
});
let result = BrokerConnection::connect(&addr_str, config).await;
assert!(
result.is_err(),
"Connection should fail when provider times out"
);
let err = match result {
Err(e) => e.to_string(),
Ok(_) => panic!("Expected error"),
};
assert!(
err.contains("timed out") || err.contains("timeout"),
"Error should mention timeout: {err}"
);
}
#[tokio::test]
async fn test_no_sasl_handshake_without_auth() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let addr_str = addr.to_string();
let mock_handle = tokio::spawn(async move {
use bytes::BufMut;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut stream, _) = listener.accept().await.unwrap();
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await.unwrap();
let len = i32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
stream.read_exact(&mut body).await.unwrap();
let api_key = i16::from_be_bytes(body[0..2].try_into().unwrap());
let api_version = i16::from_be_bytes(body[2..4].try_into().unwrap());
let correlation_id = i32::from_be_bytes(body[4..8].try_into().unwrap());
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_slice(&mock_api_versions_body(api_version));
let len = resp.len() as i32;
stream.write_all(&len.to_be_bytes()).await.unwrap();
stream.write_all(&resp).await.unwrap();
stream.flush().await.unwrap();
api_key
});
let config = ConnectionConfig::builder()
.client_id("test-client")
.build()
.unwrap();
let conn = BrokerConnection::connect(&addr_str, config).await;
assert!(conn.is_ok());
let api_key = mock_handle.await.unwrap();
assert_eq!(
api_key, 18,
"First request without auth should be ApiVersions (18), not SaslHandshake (17)"
);
conn.unwrap().close().await;
}
#[tokio::test]
async fn test_sasl_handshake_failure_rejects_connection() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let addr_str = addr.to_string();
tokio::spawn(async move {
use bytes::BufMut;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut stream, _) = listener.accept().await.unwrap();
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await.unwrap();
let len = i32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
stream.read_exact(&mut body).await.unwrap();
let correlation_id = i32::from_be_bytes(body[4..8].try_into().unwrap());
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_i16(33); resp.put_i32(1); let mech = b"GSSAPI";
resp.put_i16(mech.len() as i16);
resp.put_slice(mech);
let len = resp.len() as i32;
stream.write_all(&len.to_be_bytes()).await.unwrap();
stream.write_all(&resp).await.unwrap();
stream.flush().await.unwrap();
});
let config = ConnectionConfig::builder()
.client_id("test-client")
.auth(crate::auth::AuthConfig::sasl_plain("user", "pass").unwrap())
.build()
.unwrap();
let result = BrokerConnection::connect(&addr_str, config).await;
assert!(
result.is_err(),
"Connection should fail when SASL handshake is rejected"
);
let err = match result {
Err(e) => e,
Ok(_) => panic!("Expected error"),
};
assert!(
err.to_string().contains("SASL handshake failed"),
"Error should mention SASL handshake failure: {err}"
);
}
#[tokio::test]
async fn test_sasl_auth_failure_rejects_connection() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let addr_str = addr.to_string();
tokio::spawn(async move {
use bytes::BufMut;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut stream, _) = listener.accept().await.unwrap();
async fn read_frame(stream: &mut tokio::net::TcpStream) -> Vec<u8> {
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await.unwrap();
let len = i32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
stream.read_exact(&mut body).await.unwrap();
body
}
let req = read_frame(&mut stream).await;
let correlation_id = i32::from_be_bytes(req[4..8].try_into().unwrap());
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_i16(0); resp.put_i32(1);
let mech = b"PLAIN";
resp.put_i16(mech.len() as i16);
resp.put_slice(mech);
let len = resp.len() as i32;
stream.write_all(&len.to_be_bytes()).await.unwrap();
stream.write_all(&resp).await.unwrap();
stream.flush().await.unwrap();
let req = read_frame(&mut stream).await;
let correlation_id = i32::from_be_bytes(req[4..8].try_into().unwrap());
let mut resp = BytesMut::new();
resp.put_i32(correlation_id);
resp.put_i16(58); let msg = b"Authentication failed";
resp.put_i16(msg.len() as i16);
resp.put_slice(msg);
resp.put_i32(0); resp.put_i64(0); let len = resp.len() as i32;
stream.write_all(&len.to_be_bytes()).await.unwrap();
stream.write_all(&resp).await.unwrap();
stream.flush().await.unwrap();
});
let config = ConnectionConfig::builder()
.client_id("test-client")
.auth(crate::auth::AuthConfig::sasl_plain("user", "wrongpass").unwrap())
.build()
.unwrap();
let result = BrokerConnection::connect(&addr_str, config).await;
assert!(
result.is_err(),
"Connection should fail when authentication is rejected"
);
let err = match result {
Err(e) => e,
Ok(_) => panic!("Expected error"),
};
assert!(
err.to_string().contains("authentication failed")
|| err.to_string().contains("Authentication failed"),
"Error should mention auth failure: {err}"
);
}
#[test]
fn test_connection_config_socket_buffer_sizes() {
let mut config = ConnectionConfig::default();
assert!(config.send_buffer_size.is_none());
assert!(config.recv_buffer_size.is_none());
config.send_buffer_size = Some(1024 * 1024);
config.recv_buffer_size = Some(512 * 1024);
assert_eq!(config.send_buffer_size, Some(1024 * 1024));
assert_eq!(config.recv_buffer_size, Some(512 * 1024));
}
#[tokio::test]
async fn test_connection_invalid_address_format() {
let config = ConnectionConfig::default();
let result = BrokerConnection::connect("not-a-valid-address", config).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_read_framed_response_rejects_negative_length() {
let data: [u8; 4] = (-1i32).to_be_bytes();
let mut cursor = std::io::Cursor::new(data);
let result =
BrokerConnection::read_framed_response(&mut cursor, MAX_SASL_FRAME_BYTES).await;
assert!(result.is_err(), "negative frame length should be rejected");
let err_msg = format!("{}", result.unwrap_err());
assert!(
err_msg.contains("response length: -1"),
"error should show negative value: {err_msg}"
);
}
#[tokio::test]
async fn test_read_framed_response_rejects_zero_length() {
let data: [u8; 4] = 0i32.to_be_bytes();
let mut cursor = std::io::Cursor::new(data);
let result =
BrokerConnection::read_framed_response(&mut cursor, crate::protocol::MAX_MESSAGE_SIZE)
.await;
assert!(result.is_err(), "zero frame length should be rejected");
}
#[tokio::test]
async fn test_connection_loop_enforces_configured_max_response_size() {
let (client, mut server) = tokio::io::duplex(256);
let (reader, writer) = tokio::io::split(client);
let (_high_tx, high_rx) = mpsc::channel(4);
let (normal_tx, normal_rx) = mpsc::channel(4);
let stats = Arc::new(ConnectionStats::default());
let metrics = Arc::new(ConnectionMetrics::default());
let loop_task = tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
high_priority_rx: high_rx,
normal_priority_rx: normal_rx,
request_timeout: Duration::from_secs(30),
stats,
metrics,
max_response_size: 16,
max_in_flight_requests: 256,
max_high_priority_bypasses: 4,
},
));
let (response_tx, response_rx) = oneshot::channel();
normal_tx
.send(ConnectionCommand::Request {
data: Bytes::from_static(b"ping"),
correlation_id: 7,
api_key: ApiKey::Metadata,
api_version: 0,
response_tx,
timeout: Duration::from_secs(30),
permit: test_permit(),
})
.await
.unwrap();
let mut request = [0u8; 4];
server.read_exact(&mut request).await.unwrap();
assert_eq!(&request, b"ping");
server.write_all(&(32i32).to_be_bytes()).await.unwrap();
server.write_all(&[0u8; 32]).await.unwrap();
server.flush().await.unwrap();
let err = response_rx.await.unwrap().unwrap_err();
assert!(
err.to_string()
.contains("message size 32 exceeds maximum 16"),
"pending request should receive the configured frame-limit error: {err}"
);
let loop_err = loop_task.await.unwrap().unwrap_err();
assert!(
loop_err
.to_string()
.contains("message size 32 exceeds maximum 16"),
"connection loop should stop on oversized steady-state frames: {loop_err}"
);
}
#[test]
fn test_connection_config_default_max_response_size() {
let config = ConnectionConfig::default();
assert_eq!(
config.max_response_size,
100 * 1024 * 1024,
"default max_response_size should be MAX_MESSAGE_SIZE (100 MB)"
);
assert_eq!(
config.max_response_size,
crate::protocol::MAX_MESSAGE_SIZE,
"default max_response_size should equal protocol::MAX_MESSAGE_SIZE"
);
}
#[tokio::test]
async fn test_connection_loop_yields_to_normal_priority_after_bypass_budget() {
use tokio::io::AsyncReadExt;
let (client, mut server) = tokio::io::duplex(4096);
let (reader, writer) = tokio::io::split(client);
let (high_tx, high_rx) = mpsc::channel(16);
let (normal_tx, normal_rx) = mpsc::channel(16);
let stats = Arc::new(ConnectionStats::default());
let metrics = Arc::new(ConnectionMetrics::default());
for index in 0..8 {
let (response_tx, _response_rx) = oneshot::channel();
high_tx
.try_send(ConnectionCommand::Request {
data: Bytes::copy_from_slice(format!("H{index:03}").as_bytes()),
correlation_id: index + 1,
api_key: ApiKey::Heartbeat,
api_version: 0,
response_tx,
timeout: Duration::from_secs(30),
permit: test_permit(),
})
.unwrap();
}
let (normal_response_tx, _normal_response_rx) = oneshot::channel();
normal_tx
.try_send(ConnectionCommand::Request {
data: Bytes::from_static(b"N000"),
correlation_id: 100,
api_key: ApiKey::Produce,
api_version: 0,
response_tx: normal_response_tx,
timeout: Duration::from_secs(30),
permit: test_permit(),
})
.unwrap();
let loop_task = tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
high_priority_rx: high_rx,
normal_priority_rx: normal_rx,
request_timeout: Duration::from_secs(30),
stats: stats.clone(),
metrics: metrics.clone(),
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 32,
max_high_priority_bypasses: 4,
},
));
let mut writes = Vec::new();
for _ in 0..5 {
let mut frame = [0u8; 4];
server.read_exact(&mut frame).await.unwrap();
writes.push(String::from_utf8(frame.to_vec()).unwrap());
}
assert_eq!(writes[0], "H000");
assert_eq!(writes[1], "H001");
assert_eq!(writes[2], "H002");
assert_eq!(writes[3], "H003");
assert_eq!(writes[4], "N000");
assert_eq!(stats.bypass_yield_count(), 1);
assert_eq!(metrics.snapshot().high_priority_bypass_yields, 1);
loop_task.abort();
}
#[tokio::test]
async fn test_connection_loop_rejects_correlation_id_collision() {
use tokio::io::AsyncReadExt;
let (client, mut server) = tokio::io::duplex(4096);
let (reader, writer) = tokio::io::split(client);
let (_high_tx, high_rx) = mpsc::channel(4);
let (normal_tx, normal_rx) = mpsc::channel(4);
let stats = Arc::new(ConnectionStats::default());
let metrics = Arc::new(ConnectionMetrics::default());
let loop_task = tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
high_priority_rx: high_rx,
normal_priority_rx: normal_rx,
request_timeout: Duration::from_secs(30),
stats,
metrics,
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 32,
max_high_priority_bypasses: 4,
},
));
let (first_response_tx, first_response_rx) = oneshot::channel();
normal_tx
.send(ConnectionCommand::Request {
data: Bytes::from_static(b"req1"),
correlation_id: 77,
api_key: ApiKey::Metadata,
api_version: 0,
response_tx: first_response_tx,
timeout: Duration::from_secs(30),
permit: test_permit(),
})
.await
.unwrap();
let mut first_write = [0u8; 4];
server.read_exact(&mut first_write).await.unwrap();
assert_eq!(&first_write, b"req1");
let (second_response_tx, second_response_rx) = oneshot::channel();
normal_tx
.send(ConnectionCommand::Request {
data: Bytes::from_static(b"req2"),
correlation_id: 77,
api_key: ApiKey::Metadata,
api_version: 0,
response_tx: second_response_tx,
timeout: Duration::from_secs(30),
permit: test_permit(),
})
.await
.unwrap();
let second_err = second_response_rx.await.unwrap().unwrap_err();
assert!(second_err.to_string().contains("correlation ID collision"));
let first_err = first_response_rx.await.unwrap().unwrap_err();
assert!(first_err.to_string().contains("correlation ID collision"));
let loop_err = loop_task.await.unwrap().unwrap_err();
assert!(loop_err.to_string().contains("correlation ID collision"));
}
#[test]
fn test_connection_config_builder_max_response_size() {
let config = ConnectionConfig::builder()
.max_response_size(50 * 1024 * 1024)
.build()
.unwrap();
assert_eq!(
config.max_response_size,
50 * 1024 * 1024,
"max_response_size should be settable via builder"
);
}
#[test]
fn test_connection_config_builder_max_response_size_minimum() {
let config = ConnectionConfig::builder()
.max_response_size(100)
.build()
.unwrap();
assert_eq!(
config.max_response_size, 1024,
"max_response_size should be clamped to minimum of 1024 bytes"
);
let config_zero = ConnectionConfig::builder()
.max_response_size(0)
.build()
.unwrap();
assert_eq!(
config_zero.max_response_size, 1024,
"max_response_size(0) should clamp to 1024"
);
}
#[tokio::test]
async fn test_connect_resolves_hostname() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let hostname_addr = format!("localhost:{port}");
let config = ConnectionConfig::builder()
.connect_timeout(Duration::from_secs(2))
.request_timeout(Duration::from_secs(2))
.build()
.unwrap();
let result = BrokerConnection::connect(&hostname_addr, config).await;
match result {
Ok(_) => {} Err(err) => {
let err_msg = format!("{err}");
assert!(
!err_msg.contains("invalid address"),
"should not fail on address resolution, got: {err_msg}"
);
}
}
}
#[tokio::test]
async fn test_connect_dns_failure_is_retriable() {
let config = ConnectionConfig::builder()
.connect_timeout(Duration::from_secs(5))
.build()
.unwrap();
let result =
BrokerConnection::connect("this-host-does-not-exist.invalid:9092", config).await;
match result {
Ok(_) => panic!("connect to non-existent host should fail"),
Err(err) => {
assert!(
err.is_retriable(),
"DNS resolution failure should be retriable (Network), got: {err}"
);
}
}
}
#[cfg(feature = "socks5")]
#[test]
fn test_proxy_config_new() {
let proxy = ProxyConfig::new("proxy.example.com:1080");
assert_eq!(proxy.address(), "proxy.example.com:1080");
assert!(proxy.credentials().is_none());
}
#[cfg(feature = "socks5")]
#[test]
fn test_proxy_config_with_credentials() {
let proxy = ProxyConfig::with_credentials("proxy.example.com:1080", "user", "s3cret");
assert_eq!(proxy.address(), "proxy.example.com:1080");
let creds = proxy.credentials().expect("should have credentials");
assert_eq!(creds.username(), "user");
assert_eq!(creds.password(), "s3cret");
}
#[cfg(feature = "socks5")]
#[test]
fn test_proxy_config_debug_redacts_credentials() {
let proxy = ProxyConfig::with_credentials("proxy.example.com:1080", "admin", "hunter2");
let debug_str = format!("{proxy:?}");
assert!(
debug_str.contains("proxy.example.com:1080"),
"Debug should contain the address"
);
assert!(
!debug_str.contains("hunter2"),
"Debug must NOT contain the password"
);
assert!(
debug_str.contains("[REDACTED]"),
"Debug should show [REDACTED] for credentials"
);
}
#[cfg(feature = "socks5")]
#[test]
fn test_proxy_credentials_debug_redacts() {
let proxy = ProxyConfig::with_credentials("proxy.example.com:1080", "user", "password123");
let creds = proxy.credentials().expect("should have credentials");
let debug_str = format!("{creds:?}");
assert!(
!debug_str.contains("password123"),
"Debug must NOT contain the password"
);
assert!(
debug_str.contains("[REDACTED]"),
"Debug should show [REDACTED]"
);
}
#[cfg(feature = "socks5")]
#[test]
fn test_connection_config_builder_with_proxy() {
let proxy = ProxyConfig::new("socks5.internal:1080");
let config = ConnectionConfig::builder()
.client_id("proxy-test")
.proxy(proxy)
.build()
.unwrap();
assert!(config.proxy.is_some());
assert_eq!(
config.proxy.as_ref().unwrap().address(),
"socks5.internal:1080"
);
}
#[cfg(feature = "socks5")]
#[tokio::test]
async fn test_connect_via_proxy_dns_failure_is_retriable() {
let proxy = ProxyConfig::new("this-proxy-does-not-exist.invalid:1080");
let config = ConnectionConfig::builder()
.connect_timeout(Duration::from_secs(5))
.proxy(proxy)
.build()
.unwrap();
let result = BrokerConnection::connect("broker:9092", config).await;
match result {
Ok(_) => panic!("connect through non-existent proxy should fail"),
Err(err) => {
assert!(
err.is_retriable(),
"proxy DNS failure should be retriable (Network or Timeout), got: {err}"
);
}
}
}
#[cfg(feature = "socks5")]
#[tokio::test]
async fn test_connect_via_proxy_stalled_handshake_times_out() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy = ProxyConfig::new(listener.local_addr().unwrap().to_string());
let config = ConnectionConfig::builder()
.connect_timeout(Duration::from_millis(75))
.proxy(proxy.clone())
.build()
.unwrap();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let mock_proxy = tokio::spawn(async move {
let (_stream, _) = listener.accept().await.unwrap();
let _ = shutdown_rx.await;
});
let started_at = Instant::now();
let err = BrokerConnection::connect_via_proxy("broker.internal:9092", &proxy, &config)
.await
.unwrap_err();
assert!(matches!(err, KrafkaError::Timeout { .. }));
assert!(
err.to_string().contains("SOCKS5 proxy connection"),
"timeout should identify the proxy connect path: {err}"
);
assert!(
started_at.elapsed() < Duration::from_secs(1),
"proxy handshake timeout should respect the configured deadline"
);
let _ = shutdown_tx.send(());
mock_proxy.await.unwrap();
}
#[tokio::test]
async fn test_send_fire_and_forget_uses_reserved_correlation_id() {
let (high_priority_tx, _high_priority_rx) = mpsc::channel(1);
let (normal_priority_tx, mut normal_priority_rx) = mpsc::channel(1);
let conn = BrokerConnection {
address: "test-broker".to_string(),
config: ConnectionConfig::default(),
correlation_id_gen: Arc::new(CorrelationIdGenerator::new()),
high_priority_tx,
normal_priority_tx,
api_versions: Arc::new(parking_lot::Mutex::new(AHashMap::new())),
broker_features: Arc::new(parking_lot::Mutex::new(BrokerFeatures::default())),
alive: Arc::new(std::sync::atomic::AtomicBool::new(true)),
session_expiry: None,
stats: Arc::new(ConnectionStats::default()),
throttle_until: Arc::new(parking_lot::Mutex::new(Instant::now())),
created_at: Instant::now(),
last_used_nanos: AtomicU64::new(0),
in_flight: Arc::new(tokio::sync::Semaphore::new(1)),
};
conn.send_fire_and_forget(ApiKey::Produce, 0, |_| Ok(()))
.await
.unwrap();
let Some(ConnectionCommand::FireAndForget { data }) = normal_priority_rx.recv().await
else {
panic!("expected fire-and-forget command");
};
let frame_len = i32::from_be_bytes(data[..4].try_into().unwrap()) as usize;
assert_eq!(frame_len, data.len() - 4);
let correlation_id = i32::from_be_bytes(data[8..12].try_into().unwrap());
assert_eq!(correlation_id, NO_RESPONSE_CORRELATION_ID);
assert_eq!(conn.correlation_id_gen.next(), 1);
}
#[test]
fn test_compute_session_expiry_zero_means_no_expiry() {
assert!(
BrokerConnection::compute_session_expiry(0).is_none(),
"session_lifetime_ms = 0 should mean no expiry"
);
}
#[test]
fn test_compute_session_expiry_negative_means_no_expiry() {
assert!(
BrokerConnection::compute_session_expiry(-1).is_none(),
"negative session_lifetime_ms should mean no expiry"
);
}
#[test]
fn test_compute_session_expiry_applies_jittered_margin() {
let before = Instant::now();
let expiry = BrokerConnection::compute_session_expiry(10_000).unwrap();
let after = Instant::now();
let expected_low = before + Duration::from_millis(8_500);
let expected_high = after + Duration::from_millis(9_500);
assert!(
expiry >= expected_low && expiry <= expected_high,
"expiry should be between 8.5s and 9.5s from now (85-95% of 10s)"
);
}
#[test]
fn test_compute_session_expiry_jitter_varies() {
let results: Vec<Instant> = (0..20)
.map(|_| BrokerConnection::compute_session_expiry(100_000).unwrap())
.collect();
let first = results[0];
let any_different = results.iter().any(|r| *r != first);
assert!(
any_different,
"20 calls should produce at least one different expiry (randomised jitter)"
);
}
#[test]
fn test_compute_session_expiry_small_lifetime() {
let expiry = BrokerConnection::compute_session_expiry(100);
assert!(expiry.is_some(), "100ms lifetime should produce an expiry");
}
#[tokio::test]
async fn test_session_lifetime_tracked_from_broker() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let mock_handle = tokio::spawn(run_mock_sasl_broker_with_lifetime(
listener,
60_000,
shutdown_rx,
));
let config = ConnectionConfig::builder()
.client_id("test-client")
.auth(crate::auth::AuthConfig::sasl_plain("user", "pass").unwrap())
.build()
.unwrap();
let conn = BrokerConnection::connect(&addr, config).await.unwrap();
assert!(
conn.session_expiry().is_some(),
"session_expiry should be set when broker reports a lifetime"
);
let remaining = conn.session_expiry().unwrap() - Instant::now();
assert!(
remaining > Duration::from_secs(49) && remaining < Duration::from_secs(58),
"session expiry should be ~51-57s from now (85-95% of 60s), got {:?}",
remaining
);
assert!(
!conn.needs_reauthentication(),
"fresh connection should not need reauthentication"
);
assert!(conn.is_usable(), "fresh connection should be usable");
conn.close().await;
let _ = shutdown_tx.send(());
mock_handle.await.unwrap();
}
#[tokio::test]
async fn test_no_session_expiry_when_lifetime_zero() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap().to_string();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let mock_handle =
tokio::spawn(run_mock_sasl_broker_with_lifetime(listener, 0, shutdown_rx));
let config = ConnectionConfig::builder()
.client_id("test-client")
.auth(crate::auth::AuthConfig::sasl_plain("user", "pass").unwrap())
.build()
.unwrap();
let conn = BrokerConnection::connect(&addr, config).await.unwrap();
assert!(
conn.session_expiry().is_none(),
"session_expiry should be None when broker reports 0"
);
assert!(
!conn.needs_reauthentication(),
"should never need reauth with no session lifetime"
);
assert!(conn.is_usable());
conn.close().await;
let _ = shutdown_tx.send(());
mock_handle.await.unwrap();
}
#[test]
fn test_extract_clock_skew_secs_valid_timestamp() {
let msg = "Signature expired: 20200101T000000Z is now past";
let skew = BrokerConnection::extract_clock_skew_secs(msg);
assert!(skew < 0, "expected negative skew, got {skew}");
}
#[test]
fn test_extract_clock_skew_secs_no_timestamp() {
let msg = "some random error message";
assert_eq!(BrokerConnection::extract_clock_skew_secs(msg), 0);
}
#[test]
fn test_extract_clock_skew_secs_malformed_timestamp() {
let msg = "Signature expired: 2020XXYYT000000Z";
assert_eq!(BrokerConnection::extract_clock_skew_secs(msg), 0);
}
#[test]
fn test_extract_clock_skew_secs_invalid_calendar_date() {
assert_eq!(
BrokerConnection::extract_clock_skew_secs("foo 20201301T000000Z bar"),
0
);
assert_eq!(
BrokerConnection::extract_clock_skew_secs("foo 20200132T000000Z bar"),
0
);
assert_eq!(
BrokerConnection::extract_clock_skew_secs("foo 20200101T250000Z bar"),
0
);
}
#[test]
fn test_extract_clock_skew_secs_leap_day() {
assert_ne!(
BrokerConnection::extract_clock_skew_secs("stamp=20200229T120000Z"),
0
);
assert_eq!(
BrokerConnection::extract_clock_skew_secs("stamp=20210229T120000Z"),
0
);
}
#[test]
fn test_extract_clock_skew_secs_embedded_in_longer_message() {
let msg = "RequestTime=THIS IS TEXT; expired; server 20200101T000000Z -- request rejected";
let skew = BrokerConnection::extract_clock_skew_secs(msg);
assert!(skew < 0);
}
#[test]
fn test_extract_clock_skew_secs_multiple_timestamps_uses_last() {
let msg = "Signature not yet current: 20200101T000000Z is not yet valid, \
not before 20990101T000000Z; check your system clock";
let skew = BrokerConnection::extract_clock_skew_secs(msg);
assert!(
skew > 0,
"expected positive skew (last timestamp used), got {skew}"
);
}
#[test]
fn test_msk_iam_clock_offset_default() {
let config = ConnectionConfig::default();
assert_eq!(config.msk_iam_clock_offset_secs.load(Ordering::Relaxed), 0);
}
#[test]
fn test_parse_aws_ts_unix_epoch() {
let ts = BrokerConnection::parse_aws_ts_unix(b"19700101T000000Z");
assert_eq!(ts, Some(0));
}
#[test]
fn test_parse_aws_ts_unix_known_date() {
let ts = BrokerConnection::parse_aws_ts_unix(b"20200101T000000Z");
assert_eq!(ts, Some(1_577_836_800));
}
#[test]
fn test_parse_aws_ts_unix_leap_day_valid() {
assert!(BrokerConnection::parse_aws_ts_unix(b"20200229T000000Z").is_some());
}
#[test]
fn test_parse_aws_ts_unix_leap_day_invalid() {
assert_eq!(
BrokerConnection::parse_aws_ts_unix(b"20210229T000000Z"),
None
);
}
#[test]
fn test_parse_aws_ts_unix_invalid_month() {
assert_eq!(
BrokerConnection::parse_aws_ts_unix(b"20201301T000000Z"),
None
);
assert_eq!(
BrokerConnection::parse_aws_ts_unix(b"20200001T000000Z"),
None
);
}
#[test]
fn test_parse_aws_ts_unix_invalid_hour() {
assert_eq!(
BrokerConnection::parse_aws_ts_unix(b"20200101T250000Z"),
None
);
}
#[test]
fn test_parse_aws_ts_unix_non_digit_chars() {
assert_eq!(
BrokerConnection::parse_aws_ts_unix(b"2020XXYYT000000Z"),
None
);
}
#[test]
fn test_msk_iam_clock_offset_clamps_to_sigv4_window() {
assert_eq!(BrokerConnection::clamp_msk_iam_clock_offset_secs(450), 300);
assert_eq!(
BrokerConnection::clamp_msk_iam_clock_offset_secs(-450),
-300
);
assert_eq!(BrokerConnection::clamp_msk_iam_clock_offset_secs(120), 120);
}
#[tokio::test]
async fn test_request_times_out_when_no_response() {
let (client, _server) = tokio::io::duplex(4096);
let (reader, writer) = tokio::io::split(client);
let (_high_tx, high_rx) = mpsc::channel(4);
let (normal_tx, normal_rx) = mpsc::channel(4);
let stats = Arc::new(ConnectionStats::default());
let metrics = Arc::new(ConnectionMetrics::default());
let request_timeout = Duration::from_millis(50);
tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
high_priority_rx: high_rx,
normal_priority_rx: normal_rx,
request_timeout,
stats,
metrics,
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 256,
max_high_priority_bypasses: 4,
},
));
let (response_tx, response_rx) = oneshot::channel();
normal_tx
.send(ConnectionCommand::Request {
data: Bytes::from_static(b"test"),
correlation_id: 42,
api_key: ApiKey::Produce,
api_version: 0,
response_tx,
timeout: request_timeout,
permit: test_permit(),
})
.await
.unwrap();
let err = response_rx.await.unwrap().unwrap_err();
assert!(
err.to_string().contains("timed out"),
"expected timeout error, got: {err}"
);
}
#[tokio::test]
async fn test_response_cancels_timeout() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (client, mut server) = tokio::io::duplex(4096);
let (reader, writer) = tokio::io::split(client);
let (_high_tx, high_rx) = mpsc::channel(4);
let (normal_tx, normal_rx) = mpsc::channel(4);
let stats = Arc::new(ConnectionStats::default());
let metrics = Arc::new(ConnectionMetrics::default());
let request_timeout = Duration::from_secs(5);
let correlation_id: i32 = 99;
tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
high_priority_rx: high_rx,
normal_priority_rx: normal_rx,
request_timeout,
stats,
metrics,
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 256,
max_high_priority_bypasses: 4,
},
));
let (response_tx, response_rx) = oneshot::channel();
normal_tx
.send(ConnectionCommand::Request {
data: Bytes::from_static(b"test"),
correlation_id,
api_key: ApiKey::Produce,
api_version: 0,
response_tx,
timeout: request_timeout,
permit: test_permit(),
})
.await
.unwrap();
let mut buf = [0u8; 4];
server.read_exact(&mut buf).await.unwrap();
let body = correlation_id.to_be_bytes();
server.write_all(&(4i32).to_be_bytes()).await.unwrap();
server.write_all(&body).await.unwrap();
server.flush().await.unwrap();
let result = response_rx.await.unwrap();
assert!(
result.is_ok(),
"expected successful response before timeout, got: {:?}",
result.unwrap_err()
);
}
#[tokio::test]
async fn test_per_request_timeout_outlives_connection_request_timeout() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (client, mut server) = tokio::io::duplex(4096);
let (reader, writer) = tokio::io::split(client);
let (_high_tx, high_rx) = mpsc::channel(4);
let (normal_tx, normal_rx) = mpsc::channel(4);
let stats = Arc::new(ConnectionStats::default());
let metrics = Arc::new(ConnectionMetrics::default());
let correlation_id: i32 = 4242;
tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
high_priority_rx: high_rx,
normal_priority_rx: normal_rx,
request_timeout: Duration::from_millis(150),
stats,
metrics,
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 256,
max_high_priority_bypasses: 4,
},
));
let (response_tx, response_rx) = oneshot::channel();
normal_tx
.send(ConnectionCommand::Request {
data: Bytes::from_static(b"join"),
correlation_id,
api_key: ApiKey::JoinGroup,
api_version: 0,
response_tx,
timeout: Duration::from_secs(5),
permit: test_permit(),
})
.await
.unwrap();
let mut buf = [0u8; 4];
server.read_exact(&mut buf).await.unwrap();
tokio::time::sleep(Duration::from_millis(600)).await;
server.write_all(&(4i32).to_be_bytes()).await.unwrap();
server
.write_all(&correlation_id.to_be_bytes())
.await
.unwrap();
server.flush().await.unwrap();
let result = response_rx.await.unwrap();
assert!(
result.is_ok(),
"a request with its own longer budget must not be expired at the \
connection's request_timeout, got: {:?}",
result.unwrap_err()
);
}
#[test]
fn test_connection_closed_error_is_retriable() {
let err = connection_closed_error();
assert!(
err.is_retriable(),
"connection loss must be retriable, got: {err:?}"
);
assert!(matches!(err, KrafkaError::Network(_)));
}
#[test]
fn test_in_flight_cap_error_is_retriable() {
let err = KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::WouldBlock,
"max in-flight requests (10) reached; retry",
));
assert!(err.is_retriable());
}
#[tokio::test]
async fn test_in_flight_permits_bound_concurrency() {
let sem = Arc::new(tokio::sync::Semaphore::new(2));
let a = sem.clone().try_acquire_owned().unwrap();
let b = sem.clone().try_acquire_owned().unwrap();
assert!(
sem.clone().try_acquire_owned().is_err(),
"third acquire must block once the cap is reached"
);
drop(a);
assert!(
sem.clone().try_acquire_owned().is_ok(),
"releasing a permit must admit the next request"
);
drop(b);
}
#[tokio::test]
async fn test_read_framed_response_rejects_oversized_sasl_frame() {
let declared = (MAX_SASL_FRAME_BYTES as i32) + 1;
let mut cursor = std::io::Cursor::new(declared.to_be_bytes().to_vec());
let err = BrokerConnection::read_framed_response(&mut cursor, MAX_SASL_FRAME_BYTES)
.await
.unwrap_err();
assert!(
err.to_string().contains("pre-authentication"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn test_sasl_frame_cap_is_far_below_max_response_size() {
const { assert!(MAX_SASL_FRAME_BYTES < crate::protocol::MAX_MESSAGE_SIZE / 1000) };
}
#[tokio::test]
async fn test_read_framed_response_does_not_preallocate_declared_length() {
let declared = (MAX_SASL_FRAME_BYTES as i32) - 1;
let mut data = declared.to_be_bytes().to_vec();
data.extend_from_slice(b"only a few bytes");
let mut cursor = std::io::Cursor::new(data);
let err = BrokerConnection::read_framed_response(&mut cursor, MAX_SASL_FRAME_BYTES)
.await
.unwrap_err();
assert!(
err.to_string()
.contains("peer closed during SASL handshake"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn test_read_handshake_frame_times_out_on_silent_peer() {
let (mut client, _server) = tokio::io::duplex(64);
let deadline = tokio::time::Instant::now() + Duration::from_millis(50);
let err = BrokerConnection::read_handshake_frame(&mut client, deadline, "SaslHandshake")
.await
.unwrap_err();
assert!(matches!(err, KrafkaError::Timeout { .. }), "got: {err:?}");
}
#[tokio::test]
async fn test_read_handshake_frame_times_out_on_dribbling_peer() {
let (mut client, mut server) = tokio::io::duplex(1024);
tokio::spawn(async move {
let _ = server.write_all(&1024i32.to_be_bytes()).await;
let _ = server.write_all(b"partial").await;
std::future::pending::<()>().await;
});
let deadline = tokio::time::Instant::now() + Duration::from_millis(100);
let err = BrokerConnection::read_handshake_frame(&mut client, deadline, "SaslAuthenticate")
.await
.unwrap_err();
assert!(matches!(err, KrafkaError::Timeout { .. }), "got: {err:?}");
}
#[test]
fn test_scram_channel_binding_defaults_to_enabled() {
let auth = crate::auth::AuthConfig::sasl_scram_sha256("u", "p");
assert!(
auth.scram_channel_binding(),
"channel binding must stay on by default"
);
}
#[test]
fn test_scram_channel_binding_can_be_disabled() {
let auth =
crate::auth::AuthConfig::sasl_scram_sha256("u", "p").with_scram_channel_binding(false);
assert!(!auth.scram_channel_binding());
}
use bytes::BufMut as _;
fn body_with_leading_i32(value: i32) -> Vec<u8> {
value.to_be_bytes().to_vec()
}
#[test]
fn test_leading_throttle_is_read_for_an_api_that_reports_it_first() {
let body = body_with_leading_i32(1234);
assert_eq!(
leading_throttle_time_ms(ApiKey::CreateTopics, 4, &body),
Some(1234)
);
}
#[test]
fn test_leading_throttle_is_ignored_below_the_versions_that_carry_it() {
let body = body_with_leading_i32(3);
assert_eq!(
leading_throttle_time_ms(ApiKey::CreateTopics, 1, &body),
None
);
assert_eq!(
leading_throttle_time_ms(ApiKey::CreateTopics, 2, &body),
Some(3)
);
}
#[test]
fn test_produce_is_excluded_from_the_leading_throttle_hook() {
let body = body_with_leading_i32(50_000);
assert_eq!(leading_throttle_time_ms(ApiKey::Produce, 10, &body), None);
assert!(
ApiKey::Produce
.leading_throttle_time_min_version()
.is_none()
);
}
#[test]
fn test_apis_without_a_leading_throttle_are_excluded() {
for api_key in [
ApiKey::ApiVersions,
ApiKey::SaslHandshake,
ApiKey::SaslAuthenticate,
ApiKey::OffsetDelete,
ApiKey::CreateDelegationToken,
ApiKey::DescribeDelegationToken,
ApiKey::WriteTxnMarkers,
ApiKey::DescribeQuorum,
] {
assert!(
api_key.leading_throttle_time_min_version().is_none(),
"{api_key:?} does not report throttle_time_ms first"
);
}
}
#[test]
fn test_implausible_and_non_positive_throttles_are_ignored() {
let huge = body_with_leading_i32(MAX_HONOURED_THROTTLE_MS + 1);
assert_eq!(leading_throttle_time_ms(ApiKey::Metadata, 8, &huge), None);
let at_limit = body_with_leading_i32(MAX_HONOURED_THROTTLE_MS);
assert_eq!(
leading_throttle_time_ms(ApiKey::Metadata, 8, &at_limit),
Some(MAX_HONOURED_THROTTLE_MS)
);
for value in [0, -1, i32::MIN] {
assert_eq!(
leading_throttle_time_ms(ApiKey::Metadata, 8, &body_with_leading_i32(value)),
None
);
}
}
#[test]
fn test_a_truncated_body_is_not_peeked() {
assert_eq!(
leading_throttle_time_ms(ApiKey::Metadata, 8, &[0, 0, 1]),
None
);
assert_eq!(leading_throttle_time_ms(ApiKey::Metadata, 8, &[]), None);
}
#[test]
fn test_the_peeked_value_agrees_with_the_real_decoders() {
use crate::protocol::VersionedDecode;
let mut body = BytesMut::new();
body.put_i32(777);
body.put_i32(0);
let body = body.freeze();
assert_eq!(
leading_throttle_time_ms(ApiKey::CreateTopics, 4, &body),
Some(777)
);
let decoded =
crate::protocol::CreateTopicsResponse::decode_versioned(4, &mut body.clone()).unwrap();
assert_eq!(decoded.throttle_time_ms, 777);
let mut body = BytesMut::new();
body.put_i32(4242);
body.put_u8(1); body.put_u8(0); let body = body.freeze();
assert_eq!(
leading_throttle_time_ms(ApiKey::DeleteGroups, 2, &body),
Some(4242)
);
let decoded =
crate::protocol::DeleteGroupsResponse::decode_versioned(2, &mut body.clone()).unwrap();
assert_eq!(decoded.throttle_time_ms, 4242);
}
#[test]
fn test_a_short_request_timeout_is_accepted_once_connect_timeout_is_lowered() {
let config = ConnectionConfig::builder()
.request_timeout(Duration::from_secs(2))
.connect_timeout(Duration::from_secs(2))
.build()
.expect("a 2 s request timeout is reachable by lowering connect_timeout");
assert_eq!(config.request_timeout, Duration::from_secs(2));
assert_eq!(config.connect_timeout, Duration::from_secs(2));
}
#[test]
fn test_a_request_timeout_below_the_default_connect_timeout_is_rejected() {
let err = ConnectionConfig::builder()
.request_timeout(Duration::from_secs(2))
.build()
.expect_err("request_timeout below connect_timeout must not build");
let message = err.to_string();
assert!(
message.contains("connect_timeout"),
"the error must name the setter to change: {message}"
);
}
#[test]
fn test_the_default_connect_timeout_constant_is_what_the_builder_uses() {
let config = ConnectionConfig::default();
assert_eq!(config.connect_timeout, DEFAULT_CONNECT_TIMEOUT);
}
}