use ahash::AHashMap;
use futures::FutureExt;
use std::sync::Arc;
use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::time::{Duration, SystemTime};
use tokio::time::Instant;
use arc_swap::ArcSwap;
use bytes::{Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::sync::{mpsc, oneshot, watch};
use tokio::time::timeout;
use tokio_rustls::TlsConnector;
use tokio_util::time::{DelayQueue, delay_queue};
use tracing::{debug, error, info, trace, warn};
use crate::auth::AuthConfig;
use crate::auth::msk_iam::MAX_SIGV4_CLOCK_SKEW_SECS;
use crate::auth::tls::{build_tls_connector, connect_tls};
use crate::error::{ErrorCode, KrafkaError, ProtocolErrorKind, Result};
use crate::metrics::ConnectionRecorder;
struct ConnectionLoopParams {
address: String,
request_rx: mpsc::Receiver<ConnectionCommand>,
close_rx: watch::Receiver<CloseMode>,
throttle_until: Arc<parking_lot::Mutex<Instant>>,
metrics: Arc<ConnectionRecorder>,
max_response_size: usize,
max_in_flight_requests: usize,
request_timeout: Duration,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
enum CloseMode {
Open,
WhenIdle,
Now,
}
#[derive(Clone)]
pub struct ProxyConfig {
pub(super) address: String,
pub(super) credentials: Option<ProxyCredentials>,
}
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
}
#[cfg(test)]
fn credentials(&self) -> Option<&ProxyCredentials> {
self.credentials.as_ref()
}
}
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()
}
}
#[derive(Clone, zeroize::ZeroizeOnDrop)]
pub struct ProxyCredentials {
username: zeroize::Zeroizing<String>,
password: zeroize::Zeroizing<String>,
}
impl ProxyCredentials {
#[inline]
pub fn username(&self) -> &str {
&self.username
}
#[inline]
pub fn password(&self) -> &str {
&self.password
}
}
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};
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) max_response_size: usize,
pub(crate) max_in_flight_requests: 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<ConnectionRecorder>,
pub(crate) proxy: Option<ProxyConfig>,
#[cfg(feature = "test-broker")]
pub(crate) connector: Option<super::connector::Connector>,
}
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("max_response_size", &self.max_response_size)
.field("max_in_flight_requests", &self.max_in_flight_requests)
.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),
);
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 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(crate) fn connection_metrics(&self) -> &Arc<ConnectionRecorder> {
&self.connection_metrics
}
#[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(),
max_response_size: crate::protocol::MAX_MESSAGE_SIZE,
max_in_flight_requests: 10,
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(ConnectionRecorder::default()),
proxy: None,
#[cfg(feature = "test-broker")]
connector: 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 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 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 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
}
}
fn extend_mute(throttle_until: &parking_lot::Mutex<Instant>, throttle_time_ms: i32, address: &str) {
if throttle_time_ms <= 0 {
return;
}
let ms = throttle_time_ms.min(MAX_HONOURED_THROTTLE_MS) as u64;
let new_deadline = Instant::now() + Duration::from_millis(ms);
let mut deadline = throttle_until.lock();
if new_deadline > *deadline {
debug!(
throttle_ms = ms,
broker = %address,
"Broker throttle applied; connection muted (KIP-219)"
);
*deadline = new_deadline;
}
}
fn check_handshake_correlation(actual: i32, expected: i32, what: &str) -> Result<()> {
if actual == expected {
Ok(())
} else {
Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!(
"{what} response carries correlation_id={actual}, expected {expected}; \
the handshake stream is out of step"
),
))
}
}
fn connection_closed_error() -> KrafkaError {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"connection closed",
))
}
fn request_timeout_close_error(correlation_id: i32, timeout: Duration) -> KrafkaError {
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!(
"connection closed: request {correlation_id} timed out after {timeout:?}, and \
responses behind it cannot arrive"
),
))
}
#[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,
timeout: Duration,
_permit: tokio::sync::OwnedSemaphorePermit,
}
enum ConnectionCommand {
Request {
data: Bytes,
correlation_id: i32,
api_key: ApiKey,
api_version: i16,
response_tx: oneshot::Sender<Result<Bytes>>,
timeout: Duration,
permit: tokio::sync::OwnedSemaphorePermit,
},
FireAndForget { data: Bytes },
}
impl ConnectionCommand {
fn is_abandoned(&self) -> bool {
match self {
Self::Request { response_tx, .. } => response_tx.is_closed(),
Self::FireAndForget { .. } => false,
}
}
}
pub struct BrokerConnection {
address: String,
config: ConnectionConfig,
correlation_id_gen: Arc<CorrelationIdGenerator>,
request_tx: mpsc::Sender<ConnectionCommand>,
close_tx: watch::Sender<CloseMode>,
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>,
throttle_until: Arc<parking_lot::Mutex<Instant>>,
created_at: Instant,
last_used_nanos: AtomicU64,
in_flight: Arc<tokio::sync::Semaphore>,
}
impl BrokerConnection {
pub async fn connect(address: &str, config: ConnectionConfig) -> Result<Self> {
let stream = super::connector::dial(address, &config).await?;
debug!("Connected to broker at {address}");
let (request_tx, request_rx) = mpsc::channel(config.max_in_flight_requests.max(1));
let (close_tx, close_rx) = watch::channel(CloseMode::Open);
let throttle_until = Arc::new(parking_lot::Mutex::new(Instant::now()));
let alive = Arc::new(std::sync::atomic::AtomicBool::new(true));
let alive_clone = alive.clone();
let mut connection = Self {
address: address.to_string(),
config: config.clone(),
correlation_id_gen: Arc::new(CorrelationIdGenerator::new()),
request_tx,
close_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,
throttle_until: throttle_until.clone(),
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(),
request_rx,
close_rx,
throttle_until,
metrics: config.connection_metrics.clone(),
max_response_size: config.max_response_size,
max_in_flight_requests: config.max_in_flight_requests,
request_timeout,
};
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 = tokio::time::Instant::now();
let tls_stream = tokio::time::timeout_at(
handshake_deadline,
connect_tls(
stream.into_tcp()?,
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 session_lifetime_ms = Self::perform_sasl_handshake(
&mut tls_stream,
auth,
address,
&config.client_id,
request_timeout,
handshake_deadline,
&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,
&config.msk_iam_clock_offset_secs,
)
.await?;
connection.session_expiry = Self::effective_session_expiry(session_lifetime_ms, auth);
config.connection_metrics.record_connect();
Self::spawn_plain_connection_task(stream, loop_params, alive_clone);
} else {
config.connection_metrics.record_connect();
Self::spawn_plain_connection_task(stream, loop_params, alive_clone);
}
connection.fetch_api_versions().await?;
Ok(connection)
}
#[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,
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)?
.ok_or_else(|| KrafkaError::auth("Failed to create SASL authenticator"))?;
auth.warn_if_cleartext_credential(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)?;
Self::write_handshake_frame(stream, &encoder.take(), deadline, "SaslHandshake").await?;
let mut response_buf =
Self::read_handshake_frame(stream, deadline, "SaslHandshake").await?;
let header = ResponseHeader::decode(&mut response_buf, ApiKey::SaslHandshake, 1)?;
check_handshake_correlation(header.correlation_id, 0, "SaslHandshake")?;
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 mut correlation_id = 1;
let initial_bytes = authenticator.initial_response()?;
Self::send_sasl_authenticate(stream, &initial_bytes, client_id, correlation_id, deadline)
.await?;
let auth_response =
Self::read_sasl_authenticate_response(stream, deadline, correlation_id).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 } => {
correlation_id += 1;
let _ = Self::send_sasl_authenticate(
stream,
&ack,
client_id,
correlation_id,
deadline,
)
.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"
)));
}
correlation_id += 1;
Self::send_sasl_authenticate(
stream,
&response_bytes,
client_id,
correlation_id,
deadline,
)
.await?;
let resp =
Self::read_sasl_authenticate_response(stream, deadline, correlation_id)
.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,
correlation_id: i32,
deadline: tokio::time::Instant,
) -> 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, correlation_id)
.with_client_id(client_id);
header.encode(encoder.buffer_mut())?;
request.encode_v1(encoder.buffer_mut())?;
encoder.finish_message(pos)?;
Self::write_handshake_frame(stream, &encoder.take(), deadline, "SaslAuthenticate").await
}
async fn write_handshake_frame<S>(
stream: &mut S,
frame: &[u8],
deadline: tokio::time::Instant,
what: &str,
) -> Result<()>
where
S: AsyncWrite + Unpin,
{
tokio::time::timeout_at(deadline, async {
stream.write_all(frame).await?;
stream.flush().await
})
.await
.map_err(|_| {
KrafkaError::timeout(format!(
"timed out writing the {what} request during SASL handshake"
))
})?
.map_err(KrafkaError::network)
}
async fn read_sasl_authenticate_response<S>(
stream: &mut S,
deadline: tokio::time::Instant,
correlation_id: i32,
) -> 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)?;
check_handshake_correlation(header.correlation_id, correlation_id, "SaslAuthenticate")?;
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_plain_connection_task(
stream: super::connector::BrokerStream,
params: ConnectionLoopParams,
alive: Arc<std::sync::atomic::AtomicBool>,
) -> tokio::task::JoinHandle<()> {
match stream {
super::connector::BrokerStream::Tcp(tcp) => {
let (reader, writer) = tcp.into_split();
Self::spawn_connection_task(reader, writer, params, alive)
}
#[cfg(feature = "test-broker")]
memory @ super::connector::BrokerStream::Memory(_) => {
let (reader, writer) = tokio::io::split(memory);
Self::spawn_connection_task(reader, writer, params, alive)
}
}
}
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>(
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 request_rx,
mut close_rx,
throttle_until,
metrics,
max_response_size,
max_in_flight_requests,
request_timeout,
} = params;
let mut pending: AHashMap<i32, PendingRequest> = AHashMap::new();
let mut delay_queue: DelayQueue<i32> = DelayQueue::new();
let mut delay_keys: AHashMap<i32, delay_queue::Key> = AHashMap::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 reader = reader;
let mut decoder = Decoder::with_max_size(max_response_size);
loop {
let item = match decoder.read_frame(&mut reader).await {
Ok(Some(frame)) => Ok(frame),
Ok(None) => {
debug!("Connection closed by peer");
return;
}
Err(e) => Err(e),
};
let failed = item.is_err();
if frame_tx.send(item).await.is_err() || failed {
return;
}
}
});
let mut terminal_error: Option<KrafkaError> = None;
let mut close_mode = CloseMode::Open;
let mut parked: Option<(ConnectionCommand, Instant)> = None;
loop {
if close_mode == CloseMode::WhenIdle && pending.is_empty() && parked.is_none() {
break;
}
let mute_end = *throttle_until.lock();
tokio::select! {
biased;
changed = close_rx.changed() => {
match changed {
Ok(()) => {
close_mode = close_mode.max(*close_rx.borrow_and_update());
if close_mode == CloseMode::Now {
debug!(broker = broker_address, "Closing connection");
break;
}
}
Err(_) => break,
}
}
frame_result = frame_rx.recv() => {
match frame_result {
Some(Ok(frame)) => {
if let Err(e) = Self::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&throttle_until,
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)
}) => {
let id = expired.into_inner();
if let Some(req) = pending.remove(&id) {
delay_keys.remove(&id);
metrics.record_stalled_connection();
warn!(
correlation_id = id,
broker = broker_address,
api_key = ?req.api_key,
in_flight = pending.len(),
"Request timed out after {:?}; closing the connection",
req.timeout
);
let _ = req.response_tx.send(Err(KrafkaError::timeout(format!(
"{:?} request {id} to {broker_address} timed out after {:?}",
req.api_key, req.timeout
))));
terminal_error = Some(request_timeout_close_error(id, req.timeout));
break;
}
}
() = tokio::time::sleep_until(mute_end), if parked.is_some() => {}
cmd = request_rx.recv(), if parked.is_none() => {
match cmd {
Some(cmd) => parked = Some((cmd, Instant::now())),
None => break,
}
}
}
let Some((cmd, waiting_since)) = parked.take() else {
continue;
};
if cmd.is_abandoned() {
trace!(
broker = broker_address,
"Dropping a request whose caller has gone"
);
continue;
}
let now = Instant::now();
if *throttle_until.lock() > now {
parked = Some((cmd, waiting_since));
continue;
}
let waited = now.saturating_duration_since(waiting_since);
if waited >= Duration::from_millis(1) {
metrics.record_throttle_delay(waited);
}
if let Err(err) = Self::write_command(
&mut writer,
&mut pending,
&mut delay_queue,
&mut delay_keys,
cmd,
max_in_flight_requests,
request_timeout,
)
.await
{
terminal_error = Some(err);
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((ConnectionCommand::Request { response_tx, .. }, _)) = parked {
let _ = response_tx.send(Err(pending_error.clone()));
}
if let Some(err) = terminal_error {
return Err(err);
}
Ok(())
}
async fn write_command<W: AsyncWrite + Unpin>(
writer: &mut W,
pending: &mut AHashMap<i32, PendingRequest>,
delay_queue: &mut DelayQueue<i32>,
delay_keys: &mut AHashMap<i32, delay_queue::Key>,
cmd: ConnectionCommand,
max_in_flight_requests: usize,
request_timeout: Duration,
) -> Result<()> {
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::unavailable(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(());
}
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,
timeout: budget,
_permit: permit,
},
);
Ok(())
}
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(()),
Ok(Err(e)) => {
error!("Fire-and-forget write error: {}", e);
Err(KrafkaError::network(e))
}
Err(_) => {
error!(
"Fire-and-forget write timed out after {:?}",
request_timeout
);
Err(KrafkaError::timeout(format!(
"fire-and-forget write timed out after {request_timeout:?}"
)))
}
}
}
}
}
fn dispatch_response(
pending: &mut AHashMap<i32, PendingRequest>,
delay_queue: &mut DelayQueue<i32>,
delay_keys: &mut AHashMap<i32, delay_queue::Key>,
throttle_until: &parking_lot::Mutex<Instant>,
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();
let Some(req) = pending.remove(&correlation_id) 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()
),
));
};
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..);
if let Some(throttle_time_ms) =
leading_throttle_time_ms(req.api_key, req.api_version, &body)
{
extend_mute(throttle_until, throttle_time_ms, broker_address);
}
let _ = req.response_tx.send(Ok(body));
Ok(())
}
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(),
)));
Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!("{context}; stream desynchronized"),
))
}
}
}
async fn acquire_in_flight(&self) -> Result<tokio::sync::OwnedSemaphorePermit> {
self.in_flight
.clone()
.acquire_owned()
.await
.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> {
self.send_inner(
ApiKey::ApiVersions,
version,
self.config.request_timeout,
|buf| match version {
0..=2 => request.encode_v0(buf),
3..=4 => request.encode_v3(buf),
_ => request.encode_v5(buf),
},
)
.await
}
#[must_use]
pub fn broker_features(&self) -> BrokerFeatures {
self.broker_features.lock().clone()
}
pub fn notify_throttle(&self, throttle_time_ms: i32) {
extend_mute(&self.throttle_until, throttle_time_ms, &self.address);
}
#[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> {
self.send_inner(
api_key,
api_version,
self.config.request_timeout,
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 budget = timeout.max(self.config.request_timeout);
self.send_inner(api_key, api_version, budget, request_body)
.await
}
async fn send_inner(
&self,
api_key: ApiKey,
api_version: i16,
budget: Duration,
request_body: impl FnOnce(&mut BytesMut) -> Result<()>,
) -> Result<Bytes> {
self.mark_used();
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().await?;
let (response_tx, response_rx) = oneshot::channel();
self.request_tx
.send(ConnectionCommand::Request {
data: encoder.take(),
correlation_id,
api_key,
api_version,
response_tx,
timeout: budget,
permit,
})
.await
.map_err(|_| connection_closed_error())?;
response_rx.await.map_err(|_| connection_closed_error())?
}
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)?;
tokio::time::timeout(
self.config.request_timeout,
self.request_tx.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())?;
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 = crate::util::with_rng(rand::Rng::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 (request_tx, _) = mpsc::channel(1);
let (close_tx, _) = watch::channel(CloseMode::Open);
Self {
address: address.to_string(),
config: ConnectionConfig::default(),
correlation_id_gen: Arc::new(CorrelationIdGenerator::new()),
request_tx,
close_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,
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_stub_session_expired(address: &str) -> Self {
let mut stub = Self::test_stub_idle_for(address, Duration::ZERO);
stub.session_expiry = Some(Instant::now());
stub
}
#[cfg(test)]
pub(crate) fn test_mark_fresh(&self) {
self.mark_used();
}
#[inline]
pub fn address(&self) -> &str {
&self.address
}
#[allow(clippy::unused_async)]
pub async fn close(&self) {
self.close_now();
}
pub(crate) fn close_now(&self) {
self.request_close(CloseMode::Now);
}
pub(crate) fn close_when_idle(&self) {
self.request_close(CloseMode::WhenIdle);
}
fn request_close(&self, mode: CloseMode) {
self.close_tx.send_if_modified(|current| {
if mode > *current {
*current = mode;
true
} else {
false
}
});
}
}
#[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!(config.auth.is_none());
}
#[test]
fn test_connection_config_with_auth() {
use crate::auth::AuthConfig;
let config = ConnectionConfig::builder()
.client_id("test")
.auth(AuthConfig::sasl_plain("user", "pass"))
.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_dispatch_response_unknown_correlation_id_is_desync() {
let mut pending = AHashMap::new();
let mut delay_queue = DelayQueue::new();
let mut delay_keys = AHashMap::new();
let throttle_until = parking_lot::Mutex::new(Instant::now());
let result = BrokerConnection::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&throttle_until,
Bytes::copy_from_slice(&42i32.to_be_bytes()),
"broker-1:9092",
);
assert!(result.is_err(), "an unknown correlation id is 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,
timeout: Duration::from_secs(30),
_permit: test_permit(),
},
);
let mut delay_queue = DelayQueue::new();
let mut delay_keys = AHashMap::new();
let throttle_until = parking_lot::Mutex::new(Instant::now());
let err = BrokerConnection::dispatch_response(
&mut pending,
&mut delay_queue,
&mut delay_keys,
&throttle_until,
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",
))
.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"))
.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"))
.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() {
use tokio::io::AsyncWriteExt;
let mut t = spawn_test_loop(Duration::from_secs(30), 16);
let (cmd, rx) = test_request(7, b"ping", Duration::from_secs(30));
t.tx.send(cmd).await.unwrap();
assert_eq!(&read4(&mut t.server).await, b"ping");
t.server.write_all(&(32i32).to_be_bytes()).await.unwrap();
t.server.write_all(&[0u8; 32]).await.unwrap();
t.server.flush().await.unwrap();
let err = 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 = t.handle.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_rejects_correlation_id_collision() {
let mut t = spawn_test_loop(Duration::from_secs(30), crate::protocol::MAX_MESSAGE_SIZE);
let (cmd, first_rx) = test_request(77, b"req1", Duration::from_secs(30));
t.tx.send(cmd).await.unwrap();
assert_eq!(&read4(&mut t.server).await, b"req1");
let (cmd, second_rx) = test_request(77, b"req2", Duration::from_secs(30));
t.tx.send(cmd).await.unwrap();
let second_err = second_rx.await.unwrap().unwrap_err();
assert!(second_err.to_string().contains("correlation ID collision"));
let first_err = first_rx.await.unwrap().unwrap_err();
assert!(first_err.to_string().contains("correlation ID collision"));
let loop_err = t.handle.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}"
);
}
}
}
#[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());
}
#[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");
}
#[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"
);
}
#[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]"
);
}
#[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"
);
}
#[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}"
);
}
}
}
#[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 =
super::super::connector::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 (request_tx, mut request_rx) = mpsc::channel(1);
let (close_tx, _close_rx) = watch::channel(CloseMode::Open);
let conn = BrokerConnection {
address: "test-broker".to_string(),
config: ConnectionConfig::default(),
correlation_id_gen: Arc::new(CorrelationIdGenerator::new()),
request_tx,
close_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,
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 }) = request_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"))
.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"))
.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 request_timeout = Duration::from_millis(50);
let mut t = spawn_test_loop(request_timeout, crate::protocol::MAX_MESSAGE_SIZE);
let (cmd, rx) = test_request(42, b"test", request_timeout);
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
let err = 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() {
let request_timeout = Duration::from_millis(300);
let mut t = spawn_test_loop(request_timeout, crate::protocol::MAX_MESSAGE_SIZE);
let (cmd, rx) = test_request(99, b"test", request_timeout);
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
answer(&mut t.server, 99).await;
rx.await.unwrap().expect("response before the timeout");
tokio::time::sleep(request_timeout * 2).await;
assert!(
!t.handle.is_finished(),
"an answered request closes nothing"
);
}
#[tokio::test]
async fn test_per_request_timeout_outlives_connection_request_timeout() {
let mut t = spawn_test_loop(
Duration::from_millis(150),
crate::protocol::MAX_MESSAGE_SIZE,
);
let (cmd, rx) = test_request(4242, b"join", Duration::from_secs(5));
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
tokio::time::sleep(Duration::from_millis(600)).await;
answer(&mut t.server, 4242).await;
rx.await.unwrap().expect(
"a request with its own longer budget must not be expired at the \
connection's request_timeout",
);
}
struct TestLoop {
server: tokio::io::DuplexStream,
tx: mpsc::Sender<ConnectionCommand>,
close_tx: watch::Sender<CloseMode>,
throttle_until: Arc<parking_lot::Mutex<Instant>>,
metrics: Arc<ConnectionRecorder>,
handle: tokio::task::JoinHandle<Result<()>>,
}
fn spawn_test_loop(request_timeout: Duration, max_response_size: usize) -> TestLoop {
let (client, server) = tokio::io::duplex(4096);
let (reader, writer) = tokio::io::split(client);
let (tx, request_rx) = mpsc::channel(16);
let (close_tx, close_rx) = watch::channel(CloseMode::Open);
let throttle_until = Arc::new(parking_lot::Mutex::new(Instant::now()));
let metrics = Arc::new(ConnectionRecorder::default());
let handle = tokio::spawn(BrokerConnection::run_connection_loop(
reader,
writer,
ConnectionLoopParams {
address: "test-broker".to_string(),
request_rx,
close_rx,
throttle_until: throttle_until.clone(),
metrics: metrics.clone(),
max_response_size,
max_in_flight_requests: 32,
request_timeout,
},
));
TestLoop {
server,
tx,
close_tx,
throttle_until,
metrics,
handle,
}
}
fn test_request(
correlation_id: i32,
data: &'static [u8],
timeout: Duration,
) -> (ConnectionCommand, oneshot::Receiver<Result<Bytes>>) {
let (response_tx, response_rx) = oneshot::channel();
let cmd = ConnectionCommand::Request {
data: Bytes::from_static(data),
correlation_id,
api_key: ApiKey::Produce,
api_version: 0,
response_tx,
timeout,
permit: test_permit(),
};
(cmd, response_rx)
}
async fn answer(server: &mut tokio::io::DuplexStream, correlation_id: i32) {
use tokio::io::AsyncWriteExt;
server.write_all(&4i32.to_be_bytes()).await.unwrap();
server
.write_all(&correlation_id.to_be_bytes())
.await
.unwrap();
server.flush().await.unwrap();
}
async fn read4(server: &mut tokio::io::DuplexStream) -> [u8; 4] {
use tokio::io::AsyncReadExt;
let mut buf = [0u8; 4];
server.read_exact(&mut buf).await.unwrap();
buf
}
#[tokio::test]
async fn test_first_request_timeout_closes_the_connection() {
let request_timeout = Duration::from_millis(100);
let mut t = spawn_test_loop(request_timeout, crate::protocol::MAX_MESSAGE_SIZE);
let (cmd, first_rx) = test_request(1, b"req1", request_timeout);
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
let (cmd, second_rx) = test_request(2, b"req2", Duration::from_secs(30));
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
let (cmd, third_rx) = test_request(3, b"req3", Duration::from_secs(30));
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
let started = tokio::time::Instant::now();
let err = first_rx.await.unwrap().unwrap_err();
assert!(matches!(err, KrafkaError::Timeout { .. }), "got: {err:?}");
for rx in [second_rx, third_rx] {
let err = tokio::time::timeout(Duration::from_millis(500), rx)
.await
.expect("the others fail when the first times out")
.unwrap()
.unwrap_err();
assert!(matches!(err, KrafkaError::Network(_)), "got: {err:?}");
assert!(err.is_retriable());
}
assert!(started.elapsed() < Duration::from_millis(500));
t.handle
.await
.unwrap()
.expect_err("the loop exits with the timeout error");
assert_eq!(t.metrics.snapshot().stalled_connections, 1);
}
#[tokio::test]
async fn test_a_muted_connection_writes_nothing_until_the_mute_ends() {
use tokio::io::AsyncReadExt;
let request_timeout = Duration::from_millis(200);
let mut t = spawn_test_loop(request_timeout, crate::protocol::MAX_MESSAGE_SIZE);
extend_mute(&t.throttle_until, 400, "test-broker");
let (cmd, rx) = test_request(1, b"req1", request_timeout);
t.tx.send(cmd).await.unwrap();
let mut buf = [0u8; 4];
assert!(
tokio::time::timeout(Duration::from_millis(250), t.server.read_exact(&mut buf))
.await
.is_err(),
"nothing may be written while muted"
);
read4(&mut t.server).await;
answer(&mut t.server, 1).await;
rx.await.unwrap().expect("succeeds after the mute");
assert_eq!(t.metrics.snapshot().throttle_delays, 1);
}
#[tokio::test]
async fn test_a_response_throttle_mutes_the_connection() {
use tokio::io::AsyncWriteExt;
let mut t = spawn_test_loop(Duration::from_secs(5), crate::protocol::MAX_MESSAGE_SIZE);
let (response_tx, rx) = oneshot::channel();
t.tx.send(ConnectionCommand::Request {
data: Bytes::from_static(b"meta"),
correlation_id: 1,
api_key: ApiKey::Metadata,
api_version: 9,
response_tx,
timeout: Duration::from_secs(5),
permit: test_permit(),
})
.await
.unwrap();
read4(&mut t.server).await;
t.server.write_all(&9i32.to_be_bytes()).await.unwrap();
t.server.write_all(&1i32.to_be_bytes()).await.unwrap();
t.server.write_all(&[0u8]).await.unwrap();
t.server.write_all(&300i32.to_be_bytes()).await.unwrap();
t.server.flush().await.unwrap();
rx.await.unwrap().unwrap();
let remaining = t
.throttle_until
.lock()
.checked_duration_since(Instant::now())
.expect("muted");
assert!(remaining > Duration::from_millis(200), "{remaining:?}");
}
#[tokio::test]
async fn test_an_abandoned_request_is_not_written() {
let mut t = spawn_test_loop(Duration::from_secs(5), crate::protocol::MAX_MESSAGE_SIZE);
extend_mute(&t.throttle_until, 150, "test-broker");
let semaphore = Arc::new(tokio::sync::Semaphore::new(1));
let (response_tx, abandoned_rx) = oneshot::channel();
t.tx.send(ConnectionCommand::Request {
data: Bytes::from_static(b"req1"),
correlation_id: 1,
api_key: ApiKey::Produce,
api_version: 0,
response_tx,
timeout: Duration::from_secs(5),
permit: semaphore.clone().try_acquire_owned().unwrap(),
})
.await
.unwrap();
drop(abandoned_rx); let (cmd, kept_rx) = test_request(2, b"req2", Duration::from_secs(5));
t.tx.send(cmd).await.unwrap();
assert_eq!(&read4(&mut t.server).await, b"req2");
answer(&mut t.server, 2).await;
kept_rx.await.unwrap().unwrap();
assert_eq!(semaphore.available_permits(), 1, "the slot is released");
}
#[tokio::test]
async fn test_a_response_for_a_gone_caller_is_discarded() {
let mut t = spawn_test_loop(Duration::from_secs(5), crate::protocol::MAX_MESSAGE_SIZE);
let (cmd, gone_rx) = test_request(1, b"req1", Duration::from_secs(5));
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
drop(gone_rx);
answer(&mut t.server, 1).await;
let (cmd, rx) = test_request(2, b"req2", Duration::from_secs(5));
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
answer(&mut t.server, 2).await;
rx.await.unwrap().expect("the connection is still usable");
assert!(!t.handle.is_finished());
}
#[tokio::test]
async fn test_requests_are_written_in_submission_order() {
let mut t = spawn_test_loop(Duration::from_secs(5), crate::protocol::MAX_MESSAGE_SIZE);
let mut receivers = Vec::new();
for (id, data) in [(1, b"aaaa"), (2, b"bbbb"), (3, b"cccc")] {
let (cmd, rx) = test_request(id, data, Duration::from_secs(5));
t.tx.send(cmd).await.unwrap();
receivers.push(rx);
}
assert_eq!(&read4(&mut t.server).await, b"aaaa");
assert_eq!(&read4(&mut t.server).await, b"bbbb");
assert_eq!(&read4(&mut t.server).await, b"cccc");
}
#[tokio::test]
async fn test_close_when_idle_waits_for_pending_requests() {
let mut t = spawn_test_loop(Duration::from_secs(5), crate::protocol::MAX_MESSAGE_SIZE);
let (cmd, rx) = test_request(1, b"req1", Duration::from_secs(5));
t.tx.send(cmd).await.unwrap();
read4(&mut t.server).await;
t.close_tx.send_replace(CloseMode::WhenIdle);
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!t.handle.is_finished(), "a pending request keeps it open");
answer(&mut t.server, 1).await;
rx.await.unwrap().expect("the pending request completes");
tokio::time::timeout(Duration::from_secs(1), t.handle)
.await
.expect("closes once idle")
.unwrap()
.unwrap();
}
#[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_sasl_handshake_rejects_a_wrong_correlation_id() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let (mut client, mut server) = tokio::io::duplex(4096);
tokio::spawn(async move {
let mut len = [0u8; 4];
server.read_exact(&mut len).await.unwrap();
let mut body = vec![0u8; i32::from_be_bytes(len) as usize];
server.read_exact(&mut body).await.unwrap();
let mut resp = BytesMut::new();
resp.put_i32(5); resp.put_i16(0);
resp.put_i32(1);
resp.put_i16(5);
resp.put_slice(b"PLAIN");
server
.write_all(&(resp.len() as i32).to_be_bytes())
.await
.unwrap();
server.write_all(&resp).await.unwrap();
std::future::pending::<()>().await;
});
let auth = AuthConfig::sasl_plain("user", "pass");
let err = BrokerConnection::perform_sasl_handshake(
&mut client,
&auth,
"broker:9092",
"client",
Duration::from_secs(1),
tokio::time::Instant::now() + Duration::from_secs(1),
&Arc::new(AtomicI64::new(0)),
)
.await
.unwrap_err();
assert!(
matches!(err, KrafkaError::Protocol { .. }),
"expected a protocol error, got {err:?}"
);
assert!(err.to_string().contains("correlation_id=5"), "{err}");
}
#[tokio::test]
async fn test_sasl_handshake_write_is_bounded_by_the_deadline() {
let (mut client, _server) = tokio::io::duplex(16);
let auth = AuthConfig::sasl_plain("user", "pass");
let client_id = "c".repeat(256);
let err = tokio::time::timeout(
Duration::from_secs(2),
BrokerConnection::perform_sasl_handshake(
&mut client,
&auth,
"broker:9092",
&client_id,
Duration::from_secs(1),
tokio::time::Instant::now() + Duration::from_millis(100),
&Arc::new(AtomicI64::new(0)),
),
)
.await
.expect("the handshake deadline bounds the write")
.unwrap_err();
assert!(matches!(err, KrafkaError::Timeout { .. }), "got {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:?}");
}
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);
}
}