use std::{
collections::VecDeque,
marker::PhantomData,
net::SocketAddr,
sync::{Arc, Mutex, MutexGuard, PoisonError},
time::{Duration, Instant},
};
use bytes::{Bytes, BytesMut};
use kacrab_protocol::{
KafkaString, Result as ProtocolResult, frame,
frame::RequestFrameSpec,
generated::{
ApiKey, ApiVersionsRequestData, ApiVersionsResponseData, ErrorCode,
SaslAuthenticateRequestData, SaslAuthenticateResponseData, SaslHandshakeRequestData,
SaslHandshakeResponseData,
},
};
use tokio::{
io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufWriter, WriteHalf},
sync::{Notify, mpsc, oneshot},
};
#[cfg(feature = "gssapi")]
use super::{
GssapiAuthenticator,
kerberos::{KerberosLoginManager, kerberos_service_name},
};
use super::{
SaslClientAction, SaslClientAuthenticatorHandle, SaslClientSession, SaslMechanism,
auth::{
OAuthTokenCache, ScramExchange, oauthbearer_auth_bytes, plain_auth_bytes,
validate_sasl_extension_hooks,
},
backoff::{BackoffPolicy, BackoffState},
buffer::BufferPools,
capabilities::BrokerCapabilities,
config::{ConnectionConfig, TransportConfig},
error::{Result, WireError},
message::{RequestMessage, ResponseMessage},
pipeline::{RequestPipeline, ResponseEnvelope},
socket, tls,
};
const API_VERSIONS_HANDSHAKE_VERSION: i16 = 3;
const HANDSHAKE_CORRELATION_ID: i32 = 0;
const SASL_HANDSHAKE_CORRELATION_ID: i32 = 1;
const SASL_AUTHENTICATE_CORRELATION_ID: i32 = 2;
const SASL_HANDSHAKE_VERSION: i16 = 1;
const SASL_AUTHENTICATE_VERSION: i16 = 2;
const MIN_TIMEOUT_TICK: Duration = Duration::from_millis(1);
const MAX_TIMEOUT_TICK: Duration = Duration::from_millis(10);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BrokerEndpoint {
pub node_id: i32,
pub host: String,
pub port: u16,
pub addr: SocketAddr,
}
impl BrokerEndpoint {
#[must_use]
pub fn new(node_id: i32, addr: SocketAddr) -> Self {
Self {
node_id,
host: addr.ip().to_string(),
port: addr.port(),
addr,
}
}
#[must_use]
pub const fn from_resolved(node_id: i32, host: String, port: u16, addr: SocketAddr) -> Self {
Self {
node_id,
host,
port,
addr,
}
}
#[must_use]
pub fn host(&self) -> &str {
&self.host
}
#[must_use]
pub const fn port(&self) -> u16 {
self.port
}
}
#[derive(Debug, Clone)]
pub(crate) struct BrokerHandle {
tx: mpsc::Sender<RequestCommand>,
#[cfg(any(feature = "producer", feature = "consumer"))]
capabilities: Arc<std::sync::RwLock<Option<BrokerCapabilities>>>,
}
pub(crate) struct PendingBrokerResponse<Resp> {
rx: oneshot::Receiver<Result<ResponseEnvelope>>,
_response: PhantomData<Resp>,
}
impl<Resp> PendingBrokerResponse<Resp>
where
Resp: ResponseMessage,
{
pub(crate) async fn wait(self) -> Result<Resp> {
let envelope = self.rx.await.map_err(|_| WireError::ConnectionClosed)??;
decode_response::<Resp>(envelope)
}
}
struct RequestCommand {
api_key: ApiKey,
max_api_version: i16,
request: Box<dyn EncodableRequest>,
enqueued_at: Instant,
completion: RequestCompletion,
}
impl RequestCommand {
const fn expects_response(&self) -> bool {
self.completion.expects_response()
}
}
enum RequestCompletion {
Response(oneshot::Sender<Result<ResponseEnvelope>>),
NoResponse(oneshot::Sender<Result<()>>),
}
impl RequestCompletion {
const fn expects_response(&self) -> bool {
matches!(self, Self::Response(_))
}
fn send_error(self, error: WireError) {
match self {
Self::Response(tx) => {
let _ignored = tx.send(Err(error));
},
Self::NoResponse(tx) => {
let _ignored = tx.send(Err(error));
},
}
}
}
trait EncodableRequest: Send {
fn encoded_len(&self, version: i16) -> ProtocolResult<usize>;
fn write_body(&self, buf: &mut BytesMut, version: i16) -> ProtocolResult<()>;
}
struct OwnedRequest<Req> {
request: Req,
}
impl<Req> EncodableRequest for OwnedRequest<Req>
where
Req: RequestMessage + Send + Sync,
{
fn encoded_len(&self, version: i16) -> ProtocolResult<usize> {
self.request.encoded_len(version)
}
fn write_body(&self, buf: &mut BytesMut, version: i16) -> ProtocolResult<()> {
self.request.write_request(buf, version)?;
Ok(())
}
}
trait BrokerIo: AsyncRead + AsyncWrite + Unpin + Send {}
impl<T> BrokerIo for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
type BrokerStream = Box<dyn BrokerIo>;
impl BrokerHandle {
pub(crate) fn spawn(
endpoint: BrokerEndpoint,
client_id: String,
config: ConnectionConfig,
buffers: Arc<BufferPools>,
oauth_token_cache: Arc<tokio::sync::Mutex<OAuthTokenCache>>,
) -> Self {
let (tx, rx) = mpsc::channel(config.broker_queue_capacity);
#[cfg(feature = "gssapi")]
let kerberos_login = KerberosLoginManager::new(&config.sasl);
let capabilities = Arc::new(std::sync::RwLock::new(None));
let task = BrokerTask {
endpoint,
client_id,
config,
buffers,
oauth_token_cache,
rx,
capabilities: Arc::clone(&capabilities),
#[cfg(feature = "gssapi")]
kerberos_login,
};
let _task = tokio::spawn(task.run());
#[cfg(not(any(feature = "producer", feature = "consumer")))]
drop(capabilities);
Self {
tx,
#[cfg(any(feature = "producer", feature = "consumer"))]
capabilities,
}
}
#[cfg(any(feature = "producer", feature = "consumer"))]
pub(crate) fn negotiated_version(&self, api_key: ApiKey) -> Option<i16> {
self.capabilities
.read()
.unwrap_or_else(PoisonError::into_inner)
.as_ref()?
.version_for(api_key)
}
pub(crate) async fn send<Req, Resp>(
&self,
api_key: ApiKey,
api_version: i16,
request: &Req,
) -> Result<Resp>
where
Req: RequestMessage + Clone + Send + Sync + 'static,
Resp: ResponseMessage,
{
self.enqueue(api_key, api_version, request)?.wait().await
}
pub(crate) fn enqueue<Req, Resp>(
&self,
api_key: ApiKey,
api_version: i16,
request: &Req,
) -> Result<PendingBrokerResponse<Resp>>
where
Req: RequestMessage + Clone + Send + Sync + 'static,
Resp: ResponseMessage,
{
let (tx, rx) = oneshot::channel();
let command = RequestCommand {
api_key,
max_api_version: api_version,
request: Box::new(OwnedRequest {
request: request.clone(),
}),
enqueued_at: Instant::now(),
completion: RequestCompletion::Response(tx),
};
self.tx.try_send(command).map_err(|error| match error {
mpsc::error::TrySendError::Full(_) => WireError::Backpressure,
mpsc::error::TrySendError::Closed(_) => WireError::ConnectionClosed,
})?;
Ok(PendingBrokerResponse {
rx,
_response: PhantomData,
})
}
pub(crate) async fn send_without_response<Req>(
&self,
api_key: ApiKey,
api_version: i16,
request: &Req,
) -> Result<()>
where
Req: RequestMessage + Clone + Send + Sync + 'static,
{
let (tx, rx) = oneshot::channel();
let command = RequestCommand {
api_key,
max_api_version: api_version,
request: Box::new(OwnedRequest {
request: request.clone(),
}),
enqueued_at: Instant::now(),
completion: RequestCompletion::NoResponse(tx),
};
self.tx.try_send(command).map_err(|error| match error {
mpsc::error::TrySendError::Full(_) => WireError::Backpressure,
mpsc::error::TrySendError::Closed(_) => WireError::ConnectionClosed,
})?;
rx.await.map_err(|_| WireError::ConnectionClosed)?
}
}
struct BrokerTask {
endpoint: BrokerEndpoint,
client_id: String,
config: ConnectionConfig,
buffers: Arc<BufferPools>,
oauth_token_cache: Arc<tokio::sync::Mutex<OAuthTokenCache>>,
rx: mpsc::Receiver<RequestCommand>,
capabilities: Arc<std::sync::RwLock<Option<BrokerCapabilities>>>,
#[cfg(feature = "gssapi")]
kerberos_login: KerberosLoginManager,
}
impl BrokerTask {
async fn run(mut self) {
let mut pending = VecDeque::new();
let mut rx_open = true;
let mut backoff = reconnect_backoff_state(&self.config);
loop {
if pending.is_empty() && rx_open {
match self.rx.recv().await {
Some(command) => pending.push_back(command),
None => return,
}
}
expire_pending_commands(&mut pending, self.config.request_timeout);
if pending.is_empty() {
if rx_open {
continue;
}
return;
}
match self.connect_and_negotiate().await {
Ok((stream, capabilities)) => {
backoff.reset();
*self
.capabilities
.write()
.unwrap_or_else(PoisonError::into_inner) = Some(capabilities.clone());
if matches!(
self.serve_connection(stream, capabilities, &mut pending, &mut rx_open)
.await,
ServeOutcome::Closed
) && pending.is_empty()
{
return;
}
},
Err(error) => {
if is_fatal_setup_error(&error) {
fail_pending_setup_error(&mut pending, || clone_setup_error(&error));
if pending.is_empty() && !rx_open {
return;
}
continue;
}
expire_pending_commands(&mut pending, self.config.request_timeout);
if pending.is_empty() && !rx_open {
return;
}
let delay = match backoff.next_delay() {
Ok(delay) => delay,
Err(error) => {
if let WireError::RandomBytes(rng_error) = error {
fail_pending_setup_error(&mut pending, || {
WireError::RandomBytes(rng_error)
});
}
if pending.is_empty() && !rx_open {
return;
}
continue;
},
};
tokio::time::sleep(delay).await;
},
}
}
}
async fn serve_connection(
&mut self,
stream: BrokerStream,
capabilities: BrokerCapabilities,
pending: &mut VecDeque<RequestCommand>,
rx_open: &mut bool,
) -> ServeOutcome {
let (reader, writer) = tokio::io::split(stream);
let mut writer = BufWriter::new(writer);
let pipeline = Arc::new(Mutex::new(RequestPipeline::new(
self.config.max_in_flight_requests_per_connection,
self.config.request_timeout,
)));
let disconnect = Arc::new(Notify::new());
let _reader_task = tokio::spawn(read_response_frames(
reader,
Arc::clone(&pipeline),
Arc::clone(&disconnect),
self.config.read_buffer_capacity,
Arc::clone(&self.buffers),
));
let mut timeout_tick = tokio::time::interval(timeout_tick_duration(&self.config));
loop {
if self
.flush_pending(&mut writer, &pipeline, pending, &capabilities)
.await
.is_err()
{
lock_pipeline(&pipeline).fail_all();
return ServeOutcome::Disconnected;
}
tokio::select! {
maybe_command = self.rx.recv() => {
let Some(command) = maybe_command else {
*rx_open = false;
lock_pipeline(&pipeline).fail_all();
return ServeOutcome::Closed;
};
let admit = (!command.expects_response()
|| lock_pipeline(&pipeline).has_capacity())
&& pending.is_empty();
if admit {
pending.push_back(command);
} else {
command.completion.send_error(WireError::Backpressure);
}
},
() = disconnect.notified() => {
lock_pipeline(&pipeline).fail_all();
return ServeOutcome::Disconnected;
},
_ = timeout_tick.tick(), if !lock_pipeline(&pipeline).is_empty() || !pending.is_empty() => {
lock_pipeline(&pipeline).fail_expired();
expire_pending_commands(pending, self.config.request_timeout);
},
() = tokio::time::sleep(self.config.connections_max_idle), if lock_pipeline(&pipeline).is_empty() && pending.is_empty() => {
return ServeOutcome::Disconnected;
},
}
}
}
async fn flush_pending(
&self,
writer: &mut BufWriter<WriteHalf<BrokerStream>>,
pipeline: &Mutex<RequestPipeline>,
pending: &mut VecDeque<RequestCommand>,
capabilities: &BrokerCapabilities,
) -> Result<()> {
let mut wrote_any = false;
while pending.front().is_some_and(|command| {
!command.expects_response() || lock_pipeline(pipeline).has_capacity()
}) {
let Some(command) = pending.pop_front() else {
break;
};
if self
.write_command(writer, pipeline, command, capabilities)
.await?
{
wrote_any = true;
}
}
if wrote_any && let Err(error) = writer.flush().await {
lock_pipeline(pipeline).fail_all();
return Err(WireError::Io(error));
}
Ok(())
}
async fn write_command(
&self,
writer: &mut BufWriter<WriteHalf<BrokerStream>>,
pipeline: &Mutex<RequestPipeline>,
command: RequestCommand,
capabilities: &BrokerCapabilities,
) -> Result<bool> {
let RequestCommand {
api_key,
max_api_version,
request,
completion,
..
} = command;
let Some(api_version) = capabilities.version_for_limit(api_key, max_api_version) else {
completion.send_error(WireError::UnsupportedApiVersion(api_key));
return Ok(false);
};
let body_len = match request.encoded_len(api_version) {
Ok(body_len) => body_len,
Err(error) => {
completion.send_error(error.into());
return Ok(false);
},
};
let tx = match completion {
RequestCompletion::Response(tx) => tx,
RequestCompletion::NoResponse(tx) => {
let correlation_id = lock_pipeline(pipeline).next_correlation_id();
let spec = RequestFrameSpec {
api_key,
api_version,
correlation_id,
client_id: &self.client_id,
capacity_hint: 0,
};
let frame = match self.encode_request_frame(spec, body_len, &*request) {
Ok(frame) => frame,
Err(error) => {
let _ignored = tx.send(Err(error));
return Ok(false);
},
};
if let Err(error) = writer.write_all(&frame).await {
self.buffers.release_write(frame);
let _ignored = tx.send(Err(WireError::Io(error)));
return Err(WireError::ConnectionClosed);
}
self.buffers.release_write(frame);
let _ignored = tx.send(Ok(()));
return Ok(true);
},
};
let reserved = lock_pipeline(pipeline).reserve(api_key, api_version, tx);
let correlation_id = match reserved {
Ok(correlation_id) => correlation_id,
Err(tx) => {
let _ignored = tx.send(Err(WireError::Backpressure));
return Ok(false);
},
};
let spec = RequestFrameSpec {
api_key,
api_version,
correlation_id,
client_id: &self.client_id,
capacity_hint: 0,
};
let frame = match self.encode_request_frame(spec, body_len, &*request) {
Ok(frame) => frame,
Err(error) => {
lock_pipeline(pipeline).fail_correlation(correlation_id, error);
return Ok(false);
},
};
if let Err(error) = writer.write_all(&frame).await {
self.buffers.release_write(frame);
lock_pipeline(pipeline).fail_correlation(correlation_id, WireError::Io(error));
return Err(WireError::ConnectionClosed);
}
self.buffers.release_write(frame);
Ok(true)
}
async fn connect_and_negotiate(&self) -> Result<(BrokerStream, BrokerCapabilities)> {
let tcp = match self.config.transport {
TransportConfig::Plaintext => {
socket::resolve_and_connect(
&self.config.socket,
self.config.socket_connection_setup_timeout,
&socket::ResolveTarget {
host: &self.endpoint.host,
port: self.endpoint.port,
use_all_dns_ips: self.config.use_all_dns_ips,
fallback: self.endpoint.addr,
},
)
.await?
},
};
let mut stream: BrokerStream = if self.config.security.protocol.uses_tls() {
Box::new(tls::connect_client(tcp, &self.config.tls, &self.tls_server_name()).await?)
} else {
Box::new(tcp)
};
let capabilities = self.api_versions(&mut stream).await?;
if self.config.security.protocol.uses_sasl() {
self.sasl_authenticate(&mut stream).await?;
}
Ok((stream, capabilities))
}
fn tls_server_name(&self) -> String {
self.endpoint.host.clone()
}
async fn api_versions(&self, stream: &mut BrokerStream) -> Result<BrokerCapabilities> {
let request = ApiVersionsRequestData {
client_software_name: KafkaString::from("kacrab".to_owned()),
client_software_version: KafkaString::from(env!("CARGO_PKG_VERSION").to_owned()),
_unknown_tagged_fields: Vec::new(),
};
let api_version = API_VERSIONS_HANDSHAKE_VERSION;
let body_len = request.encoded_len(api_version)?;
let frame = self.encode_request_frame_with_body(
RequestFrameSpec {
api_key: ApiKey::ApiVersions,
api_version,
correlation_id: HANDSHAKE_CORRELATION_ID,
client_id: &self.client_id,
capacity_hint: 0,
},
body_len,
|buf| {
request.write(buf, api_version)?;
Ok(())
},
)?;
self.write_pooled_frame(stream, frame).await?;
let response_bytes =
read_frame(stream, self.config.read_buffer_capacity, &self.buffers).await?;
let mut response =
frame::decode_response_envelope(ApiKey::ApiVersions, api_version, response_bytes)?;
if response.correlation_id != HANDSHAKE_CORRELATION_ID {
return Err(WireError::CorrelationIdMismatch {
expected: HANDSHAKE_CORRELATION_ID,
actual: response.correlation_id,
});
}
let response = ApiVersionsResponseData::read(&mut response.body, api_version)?;
let error = ErrorCode::from(response.error_code);
if error.is_error() {
return Err(WireError::Kafka(error));
}
Ok(BrokerCapabilities::from_response(&response))
}
async fn sasl_authenticate(&self, stream: &mut BrokerStream) -> Result<()> {
validate_sasl_extension_hooks(&self.config.sasl)?;
if let Some(factory) = &self.config.sasl.client_authenticator_factory {
let session = SaslClientSession::new(
self.endpoint.node_id,
self.endpoint.host.clone(),
self.endpoint.port,
self.endpoint.addr,
);
let authenticator = factory.create(&session)?;
return self.sasl_custom_authenticate(stream, &authenticator).await;
}
if let Some(authenticator) = &self.config.sasl.client_authenticator {
return self.sasl_custom_authenticate(stream, authenticator).await;
}
let mechanism = self.config.sasl.mechanism.unwrap_or(SaslMechanism::Gssapi);
if mechanism == SaslMechanism::Gssapi {
#[cfg(feature = "gssapi")]
{
return self.sasl_gssapi_authenticate(stream).await;
}
#[cfg(not(feature = "gssapi"))]
{
return Err(WireError::GssapiBackendUnavailable);
}
}
self.sasl_handshake(stream, mechanism).await?;
let auth_bytes = match mechanism {
SaslMechanism::Plain => plain_auth_bytes(self.config.sasl.jaas_config.as_deref())?,
SaslMechanism::ScramSha256 | SaslMechanism::ScramSha512 => {
return self.sasl_scram_authenticate(stream, mechanism).await;
},
SaslMechanism::OAuthBearer => return self.sasl_oauthbearer_authenticate(stream).await,
SaslMechanism::Gssapi => return Err(WireError::GssapiBackendUnavailable),
};
let _response = self.sasl_authenticate_round(stream, auth_bytes).await?;
Ok(())
}
async fn sasl_oauthbearer_authenticate(&self, stream: &mut BrokerStream) -> Result<()> {
let auth_bytes =
oauthbearer_auth_bytes(&self.config.sasl, &self.config.tls, &self.oauth_token_cache)
.await?;
let response = self.sasl_authenticate_round(stream, auth_bytes).await?;
if response.auth_bytes.is_empty() {
return Ok(());
}
let error = String::from_utf8_lossy(response.auth_bytes.as_ref()).into_owned();
let _result = self
.sasl_authenticate_round(stream, Bytes::from_static(&[0x01]))
.await;
Err(WireError::SaslAuthentication(format!(
"OAUTHBEARER token rejected by broker: {error}"
)))
}
async fn sasl_custom_authenticate(
&self,
stream: &mut BrokerStream,
authenticator: &SaslClientAuthenticatorHandle,
) -> Result<()> {
self.sasl_handshake(stream, authenticator.mechanism())
.await?;
let mut action = authenticator.start()?;
loop {
let SaslClientAction::Send(auth_bytes) = action else {
return Ok(());
};
let response = self.sasl_authenticate_round(stream, auth_bytes).await?;
action = authenticator.next(response.auth_bytes.as_ref())?;
}
}
#[cfg(feature = "gssapi")]
async fn sasl_gssapi_authenticate(&self, stream: &mut BrokerStream) -> Result<()> {
let service_name = kerberos_service_name(&self.config.sasl)?;
let authenticator = SaslClientAuthenticatorHandle::new(
GssapiAuthenticator::new(service_name, self.gssapi_host_name())
.with_kerberos_login(self.kerberos_login.clone()),
);
self.sasl_custom_authenticate(stream, &authenticator).await
}
#[cfg(feature = "gssapi")]
fn gssapi_host_name(&self) -> String {
self.endpoint.host.clone()
}
async fn sasl_scram_authenticate(
&self,
stream: &mut BrokerStream,
mechanism: SaslMechanism,
) -> Result<()> {
let (exchange, client_first) =
ScramExchange::start(mechanism, self.config.sasl.jaas_config.as_deref())?;
let server_first = self.sasl_authenticate_round(stream, client_first).await?;
let (client_final, expected_signature) =
exchange.client_final(server_first.auth_bytes.as_ref())?;
let server_final = self.sasl_authenticate_round(stream, client_final).await?;
ScramExchange::verify_server_final(server_final.auth_bytes.as_ref(), &expected_signature)
}
async fn sasl_authenticate_round(
&self,
stream: &mut BrokerStream,
auth_bytes: Bytes,
) -> Result<SaslAuthenticateResponseData> {
let request = SaslAuthenticateRequestData {
auth_bytes,
_unknown_tagged_fields: Vec::new(),
};
let body_len = request.encoded_len(SASL_AUTHENTICATE_VERSION)?;
let frame = self.encode_request_frame_with_body(
RequestFrameSpec {
api_key: ApiKey::SaslAuthenticate,
api_version: SASL_AUTHENTICATE_VERSION,
correlation_id: SASL_AUTHENTICATE_CORRELATION_ID,
client_id: &self.client_id,
capacity_hint: 0,
},
body_len,
|buf| {
request.write(buf, SASL_AUTHENTICATE_VERSION)?;
Ok(())
},
)?;
self.write_pooled_frame(stream, frame).await?;
let response_bytes =
read_frame(stream, self.config.read_buffer_capacity, &self.buffers).await?;
let mut envelope = frame::decode_response_envelope(
ApiKey::SaslAuthenticate,
SASL_AUTHENTICATE_VERSION,
response_bytes,
)?;
if envelope.correlation_id != SASL_AUTHENTICATE_CORRELATION_ID {
return Err(WireError::CorrelationIdMismatch {
expected: SASL_AUTHENTICATE_CORRELATION_ID,
actual: envelope.correlation_id,
});
}
let response =
SaslAuthenticateResponseData::read(&mut envelope.body, SASL_AUTHENTICATE_VERSION)?;
let error = ErrorCode::from(response.error_code);
if error.is_error() {
let message = response
.error_message
.map_or_else(|| error.to_string(), |message| message.to_string());
return Err(WireError::SaslAuthentication(message));
}
Ok(response)
}
async fn sasl_handshake(
&self,
stream: &mut BrokerStream,
mechanism: SaslMechanism,
) -> Result<()> {
let request = SaslHandshakeRequestData {
mechanism: KafkaString::from(mechanism.as_str().to_owned()),
_unknown_tagged_fields: Vec::new(),
};
let body_len = request.encoded_len(SASL_HANDSHAKE_VERSION)?;
let frame = self.encode_request_frame_with_body(
RequestFrameSpec {
api_key: ApiKey::SaslHandshake,
api_version: SASL_HANDSHAKE_VERSION,
correlation_id: SASL_HANDSHAKE_CORRELATION_ID,
client_id: &self.client_id,
capacity_hint: 0,
},
body_len,
|buf| {
request.write(buf, SASL_HANDSHAKE_VERSION)?;
Ok(())
},
)?;
self.write_pooled_frame(stream, frame).await?;
let response_bytes =
read_frame(stream, self.config.read_buffer_capacity, &self.buffers).await?;
let mut envelope = frame::decode_response_envelope(
ApiKey::SaslHandshake,
SASL_HANDSHAKE_VERSION,
response_bytes,
)?;
if envelope.correlation_id != SASL_HANDSHAKE_CORRELATION_ID {
return Err(WireError::CorrelationIdMismatch {
expected: SASL_HANDSHAKE_CORRELATION_ID,
actual: envelope.correlation_id,
});
}
let response = SaslHandshakeResponseData::read(&mut envelope.body, SASL_HANDSHAKE_VERSION)?;
let error = ErrorCode::from(response.error_code);
if error.is_error() {
return Err(WireError::SaslHandshake(error.to_string()));
}
let supported = response
.mechanisms
.iter()
.any(|supported| supported.as_str() == mechanism.as_str());
if !supported {
return Err(WireError::UnsupportedSaslMechanism(
mechanism.as_str().to_owned(),
));
}
Ok(())
}
fn encode_request_frame(
&self,
spec: RequestFrameSpec<'_>,
body_len: usize,
request: &dyn EncodableRequest,
) -> Result<BytesMut> {
self.encode_request_frame_with_body(spec, body_len, |buf| {
request.write_body(buf, spec.api_version)
})
}
fn encode_request_frame_with_body<F>(
&self,
spec: RequestFrameSpec<'_>,
body_len: usize,
write_body: F,
) -> Result<BytesMut>
where
F: FnOnce(&mut BytesMut) -> ProtocolResult<()>,
{
let capacity_hint = frame::request_frame_capacity_hint(spec, body_len)?;
let mut frame = self.buffers.acquire_write(capacity_hint);
if let Err(error) = frame::encode_request_frame_with_buffer(
&mut frame,
RequestFrameSpec {
capacity_hint,
..spec
},
write_body,
) {
self.buffers.release_write(frame);
return Err(error.into());
}
Ok(frame)
}
async fn write_pooled_frame(&self, stream: &mut BrokerStream, frame: BytesMut) -> Result<()> {
if let Err(error) = stream.write_all(&frame).await {
self.buffers.release_write(frame);
return Err(WireError::Io(error));
}
if let Err(error) = stream.flush().await {
self.buffers.release_write(frame);
return Err(WireError::Io(error));
}
self.buffers.release_write(frame);
Ok(())
}
#[cfg(test)]
fn request_frame_capacity_hint(
&self,
api_key: ApiKey,
api_version: i16,
correlation_id: i32,
body_len: usize,
) -> Result<usize> {
Ok(frame::request_frame_capacity_hint(
RequestFrameSpec {
api_key,
api_version,
correlation_id,
client_id: &self.client_id,
capacity_hint: 0,
},
body_len,
)?)
}
}
enum ServeOutcome {
Disconnected,
Closed,
}
fn expire_pending_commands(pending: &mut VecDeque<RequestCommand>, request_timeout: Duration) {
let now = Instant::now();
let mut retained = VecDeque::with_capacity(pending.len());
while let Some(command) = pending.pop_front() {
if now.duration_since(command.enqueued_at) >= request_timeout {
command.completion.send_error(WireError::Timeout);
} else {
retained.push_back(command);
}
}
*pending = retained;
}
fn fail_pending_setup_error(
pending: &mut VecDeque<RequestCommand>,
mut error_factory: impl FnMut() -> WireError,
) {
while let Some(command) = pending.pop_front() {
command.completion.send_error(error_factory());
}
}
const fn is_fatal_setup_error(error: &WireError) -> bool {
matches!(
error,
WireError::UnsupportedTlsOption(_)
| WireError::InvalidTlsConfig(_)
| WireError::TlsHandshake(_)
| WireError::GssapiBackendUnavailable
| WireError::InvalidSaslConfig(_)
| WireError::UnsupportedSaslMechanism(_)
| WireError::SaslHandshake(_)
| WireError::SaslAuthentication(_)
| WireError::SaslServerSignatureMismatch
)
}
fn clone_setup_error(error: &WireError) -> WireError {
match error {
WireError::UnsupportedTlsOption(message) => {
WireError::UnsupportedTlsOption(message.clone())
},
WireError::InvalidTlsConfig(message) => WireError::InvalidTlsConfig(message.clone()),
WireError::TlsHandshake(message) => WireError::TlsHandshake(message.clone()),
WireError::GssapiBackendUnavailable => WireError::GssapiBackendUnavailable,
WireError::InvalidSaslConfig(message) => WireError::InvalidSaslConfig(message.clone()),
WireError::UnsupportedSaslMechanism(message) => {
WireError::UnsupportedSaslMechanism(message.clone())
},
WireError::SaslHandshake(message) => WireError::SaslHandshake(message.clone()),
WireError::SaslAuthentication(message) => WireError::SaslAuthentication(message.clone()),
WireError::SaslServerSignatureMismatch => WireError::SaslServerSignatureMismatch,
_ => WireError::ConnectionClosed,
}
}
fn reconnect_backoff_state(config: &ConnectionConfig) -> BackoffState {
BackoffState::new(
BackoffPolicy::new(
config.reconnect_backoff_initial,
config.reconnect_backoff_max,
)
.with_jitter_factor(super::backoff::DEFAULT_JITTER_FACTOR),
)
}
fn lock_pipeline(pipeline: &Mutex<RequestPipeline>) -> MutexGuard<'_, RequestPipeline> {
pipeline.lock().unwrap_or_else(PoisonError::into_inner)
}
async fn read_response_frames<R>(
mut reader: R,
pipeline: Arc<Mutex<RequestPipeline>>,
disconnect: Arc<Notify>,
read_buffer_capacity: Option<usize>,
buffers: Arc<BufferPools>,
) where
R: AsyncRead + Unpin,
{
loop {
match read_frame(&mut reader, read_buffer_capacity, &buffers).await {
Ok(frame) => lock_pipeline(&pipeline).complete_response(frame),
Err(_error) => {
lock_pipeline(&pipeline).fail_all();
disconnect.notify_one();
return;
},
}
}
}
async fn read_frame<R>(
reader: &mut R,
read_buffer_capacity: Option<usize>,
buffers: &BufferPools,
) -> Result<Bytes>
where
R: AsyncRead + Unpin,
{
let length = reader.read_i32().await?;
if !(0..=frame::MAX_FRAME_LENGTH).contains(&length) {
return Err(WireError::ConnectionClosed);
}
let length = usize::try_from(length).map_err(|_| WireError::ConnectionClosed)?;
let capacity = read_buffer_capacity.map_or(length, |capacity| length.max(capacity));
let mut payload = buffers.acquire_read(capacity);
payload.resize(length, 0);
let _bytes_read = reader.read_exact(&mut payload[..]).await?;
let frozen = payload.split_to(length).freeze();
buffers.release_read(payload);
Ok(frozen)
}
fn decode_response<Resp>(mut envelope: ResponseEnvelope) -> Result<Resp>
where
Resp: ResponseMessage,
{
Ok(Resp::read_response(
&mut envelope.body,
envelope.api_version,
)?)
}
fn timeout_tick_duration(config: &ConnectionConfig) -> Duration {
let timeout = config.request_timeout;
let half_timeout = timeout.checked_div(2).unwrap_or(timeout);
half_timeout.min(MAX_TIMEOUT_TICK).max(MIN_TIMEOUT_TICK)
}
#[cfg(test)]
mod tests {
#![allow(
clippy::expect_used,
clippy::missing_assert_message,
clippy::unwrap_used,
reason = "Unit test fixtures fail fastest with contextual unwrap/expect calls."
)]
use std::{
collections::VecDeque,
sync::{Arc, Mutex},
time::{Duration, Instant},
};
use bytes::{Bytes, BytesMut};
use kacrab_protocol::{
KafkaString, frame,
frame::RequestFrameSpec,
generated::{
ApiKey, ApiVersion, ApiVersionsRequestData, ApiVersionsResponseData, ErrorCode,
RequestHeaderData, SaslAuthenticateRequestData, SaslAuthenticateResponseData,
SaslHandshakeResponseData,
},
version::{request_header_version, response_header_version},
};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt, BufWriter},
net::TcpListener,
sync::{Notify, mpsc, oneshot},
};
#[cfg(feature = "gssapi")]
use super::KerberosLoginManager;
use super::{
BrokerCapabilities, BrokerEndpoint, BrokerHandle, BrokerStream, BrokerTask, BufferPools,
EncodableRequest, OAuthTokenCache, OwnedRequest, RequestCommand, RequestCompletion,
RequestPipeline, ResponseEnvelope, ServeOutcome, expire_pending_commands, lock_pipeline,
read_frame, read_response_frames, reconnect_backoff_state, timeout_tick_duration,
};
use crate::wire::{
ConnectionConfig, Result as WireResult, SaslClientAction, SaslClientAuthenticator,
SaslClientAuthenticatorFactory, SaslClientAuthenticatorHandle, SaslClientSession,
SaslMechanism, SecurityProtocol, WireError,
};
#[derive(Debug)]
struct StaticSaslAuthenticator {
mechanism: SaslMechanism,
payload: Bytes,
}
impl SaslClientAuthenticator for StaticSaslAuthenticator {
fn mechanism(&self) -> SaslMechanism {
self.mechanism
}
fn start(&self) -> WireResult<SaslClientAction> {
Ok(SaslClientAction::Send(self.payload.clone()))
}
fn next(&self, _challenge: &[u8]) -> WireResult<SaslClientAction> {
Ok(SaslClientAction::Complete)
}
}
#[derive(Debug)]
struct SessionPayloadFactory;
impl SaslClientAuthenticatorFactory for SessionPayloadFactory {
fn mechanism(&self) -> SaslMechanism {
SaslMechanism::Plain
}
fn create(&self, session: &SaslClientSession) -> WireResult<SaslClientAuthenticatorHandle> {
Ok(SaslClientAuthenticatorHandle::new(
StaticSaslAuthenticator {
mechanism: SaslMechanism::Plain,
payload: Bytes::from(format!(
"{}:{}:{}",
session.node_id(),
session.host(),
session.port()
)),
},
))
}
}
fn api_versions_request() -> ApiVersionsRequestData {
ApiVersionsRequestData {
client_software_name: KafkaString::from("kacrab".to_owned()),
client_software_version: KafkaString::from("0.0.1".to_owned()),
_unknown_tagged_fields: Vec::new(),
}
}
fn request_command() -> RequestCommand {
let (tx, _rx) = oneshot::channel();
request_command_with_sender(tx, Instant::now())
}
fn request_command_with_sender(
tx: oneshot::Sender<WireResult<ResponseEnvelope>>,
enqueued_at: Instant,
) -> RequestCommand {
RequestCommand {
api_key: ApiKey::ApiVersions,
max_api_version: 3,
request: Box::new(OwnedRequest {
request: api_versions_request(),
}),
enqueued_at,
completion: RequestCompletion::Response(tx),
}
}
fn api_versions_capabilities() -> BrokerCapabilities {
BrokerCapabilities::from_response(&ApiVersionsResponseData {
api_keys: vec![ApiVersion {
api_key: ApiKey::ApiVersions as i16,
min_version: 0,
max_version: 4,
_unknown_tagged_fields: Vec::new(),
}],
..ApiVersionsResponseData::default()
})
}
fn api_versions_response(correlation_id: i32, error: ErrorCode) -> BytesMut {
let mut header = BytesMut::new();
kacrab_protocol::generated::ResponseHeaderData {
correlation_id,
_unknown_tagged_fields: Vec::new(),
}
.write(
&mut header,
response_header_version(ApiKey::ApiVersions as i16, 3),
)
.expect("response header");
let mut body = BytesMut::new();
ApiVersionsResponseData {
error_code: error.code(),
..ApiVersionsResponseData::default()
}
.write(&mut body, 3)
.expect("api versions response");
frame::encode_request(&header, &body).expect("response frame")
}
fn sasl_handshake_response(correlation_id: i32, mechanism: &str) -> BytesMut {
let mut header = BytesMut::new();
kacrab_protocol::generated::ResponseHeaderData {
correlation_id,
_unknown_tagged_fields: Vec::new(),
}
.write(
&mut header,
response_header_version(ApiKey::SaslHandshake as i16, 1),
)
.expect("response header");
let mut body = BytesMut::new();
SaslHandshakeResponseData {
error_code: ErrorCode::None.code(),
mechanisms: vec![KafkaString::from(mechanism.to_owned())],
_unknown_tagged_fields: Vec::new(),
}
.write(&mut body, 1)
.expect("sasl handshake response");
frame::encode_request(&header, &body).expect("response frame")
}
fn sasl_authenticate_response(correlation_id: i32) -> BytesMut {
let mut header = BytesMut::new();
kacrab_protocol::generated::ResponseHeaderData {
correlation_id,
_unknown_tagged_fields: Vec::new(),
}
.write(
&mut header,
response_header_version(ApiKey::SaslAuthenticate as i16, 2),
)
.expect("response header");
let mut body = BytesMut::new();
SaslAuthenticateResponseData {
error_code: ErrorCode::None.code(),
error_message: None,
auth_bytes: Bytes::new(),
session_lifetime_ms: 0,
_unknown_tagged_fields: Vec::new(),
}
.write(&mut body, 2)
.expect("sasl authenticate response");
frame::encode_request(&header, &body).expect("response frame")
}
async fn broker_task_with_connected_stream()
-> (BrokerTask, tokio::net::TcpStream, tokio::net::TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind listener");
let addr = listener.local_addr().expect("listener addr");
let client = tokio::net::TcpStream::connect(addr)
.await
.expect("client connect");
let (server, _peer) = listener.accept().await.expect("server accept");
let (_tx, rx) = mpsc::channel(1);
let config = ConnectionConfig::default();
#[cfg(feature = "gssapi")]
let kerberos_login = KerberosLoginManager::new(&config.sasl);
let task = BrokerTask {
endpoint: BrokerEndpoint::new(7, addr),
client_id: "client-a".to_owned(),
config,
buffers: Arc::new(BufferPools::new(1)),
oauth_token_cache: Arc::new(tokio::sync::Mutex::new(OAuthTokenCache::default())),
rx,
capabilities: Arc::new(std::sync::RwLock::new(None)),
#[cfg(feature = "gssapi")]
kerberos_login,
};
(task, client, server)
}
#[test]
fn reconnect_backoff_state_resets_after_successful_connection() {
let config = ConnectionConfig::default()
.reconnect_backoff_initial(Duration::from_millis(5))
.reconnect_backoff_max(Duration::from_millis(20));
let mut state = reconnect_backoff_state(&config);
assert_eq!(state.next_delay_with_sample(0.5), Duration::from_millis(5));
assert_eq!(state.next_delay_with_sample(0.5), Duration::from_millis(10));
state.reset();
assert_eq!(state.next_delay_with_sample(0.5), Duration::from_millis(5));
}
#[test]
fn timeout_tick_duration_stays_between_floor_and_ceiling() {
assert_eq!(
timeout_tick_duration(&ConnectionConfig::default().request_timeout(Duration::ZERO)),
Duration::from_millis(1)
);
assert_eq!(
timeout_tick_duration(
&ConnectionConfig::default().request_timeout(Duration::from_mins(1))
),
Duration::from_millis(10)
);
}
#[tokio::test]
async fn expire_pending_commands_sends_timeout_and_retains_fresh_commands() {
let (expired_tx, expired_rx) = oneshot::channel();
let (fresh_tx, fresh_rx) = oneshot::channel();
let now = Instant::now();
let expired_at = now.checked_sub(Duration::from_mins(1)).unwrap_or(now);
let mut pending = VecDeque::from([
request_command_with_sender(expired_tx, expired_at),
request_command_with_sender(fresh_tx, now),
]);
expire_pending_commands(&mut pending, Duration::from_secs(30));
assert!(matches!(
expired_rx.await.expect("expired sender"),
Err(WireError::Timeout)
));
assert_eq!(pending.len(), 1);
drop(pending);
assert!(fresh_rx.await.is_err());
}
#[tokio::test]
async fn broker_handle_reports_full_and_closed_queue() {
let (tx, mut rx) = mpsc::channel(1);
tx.try_send(request_command()).expect("prefill queue");
let full = BrokerHandle {
tx,
capabilities: Arc::new(std::sync::RwLock::new(None)),
};
assert!(matches!(
full.send::<_, ApiVersionsResponseData>(
ApiKey::ApiVersions,
3,
&api_versions_request()
)
.await,
Err(WireError::Backpressure)
));
let _dropped = rx.recv().await;
let (tx, rx) = mpsc::channel(1);
drop(rx);
let closed = BrokerHandle {
tx,
capabilities: Arc::new(std::sync::RwLock::new(None)),
};
assert!(matches!(
closed
.send::<_, ApiVersionsResponseData>(ApiKey::ApiVersions, 3, &api_versions_request())
.await,
Err(WireError::ConnectionClosed)
));
}
#[test]
fn broker_task_encode_frame_wraps_body_with_client_header() {
let config = ConnectionConfig::default();
#[cfg(feature = "gssapi")]
let kerberos_login = KerberosLoginManager::new(&config.sasl);
let task = BrokerTask {
endpoint: BrokerEndpoint::new(7, "127.0.0.1:9092".parse().expect("socket address")),
client_id: "client-a".to_owned(),
config,
buffers: Arc::new(BufferPools::new(1)),
oauth_token_cache: Arc::new(tokio::sync::Mutex::new(OAuthTokenCache::default())),
rx: mpsc::channel(1).1,
capabilities: Arc::new(std::sync::RwLock::new(None)),
#[cfg(feature = "gssapi")]
kerberos_login,
};
let request = OwnedRequest {
request: api_versions_request(),
};
let body_len = request.encoded_len(3).expect("body length");
let expected_len = task
.request_frame_capacity_hint(ApiKey::ApiVersions, 3, 9, body_len)
.expect("frame capacity hint");
let frame = task
.encode_request_frame(
RequestFrameSpec {
api_key: ApiKey::ApiVersions,
api_version: 3,
correlation_id: 9,
client_id: &task.client_id,
capacity_hint: 0,
},
body_len,
&request,
)
.expect("encoded frame");
assert_eq!(frame.len(), expected_len);
}
#[tokio::test]
async fn api_versions_rejects_mismatched_correlation_and_broker_error() {
let (task, client, mut server) = broker_task_with_connected_stream().await;
let mut client: BrokerStream = Box::new(client);
let server_task = tokio::spawn(async move {
let _request_len = server.read_i32().await.expect("request length");
server
.write_all(&api_versions_response(99, ErrorCode::None))
.await
.expect("write response");
});
assert!(matches!(
task.api_versions(&mut client).await,
Err(WireError::CorrelationIdMismatch {
expected: 0,
actual: 99
})
));
server_task.await.expect("server task");
let (task, client, mut server) = broker_task_with_connected_stream().await;
let mut client: BrokerStream = Box::new(client);
let server_task = tokio::spawn(async move {
let _request_len = server.read_i32().await.expect("request length");
server
.write_all(&api_versions_response(0, ErrorCode::UnsupportedVersion))
.await
.expect("write response");
});
assert!(matches!(
task.api_versions(&mut client).await,
Err(WireError::Kafka(ErrorCode::UnsupportedVersion))
));
server_task.await.expect("server task");
}
#[tokio::test]
async fn sasl_authenticate_uses_native_rust_authenticator_payload() {
let (mut task, client, mut server) = broker_task_with_connected_stream().await;
task.config.security.protocol = SecurityProtocol::SaslPlaintext;
task.config.sasl = task
.config
.sasl
.clone()
.client_authenticator(StaticSaslAuthenticator {
mechanism: SaslMechanism::Plain,
payload: Bytes::from_static(b"native-hook-payload"),
});
let mut client: BrokerStream = Box::new(client);
let server_task = tokio::spawn(async move {
let _handshake_request = read_frame(&mut server, None, &BufferPools::new(1))
.await
.expect("read handshake request");
server
.write_all(&sasl_handshake_response(1, "PLAIN"))
.await
.expect("write handshake response");
let authenticate_request = read_frame(&mut server, None, &BufferPools::new(1))
.await
.expect("read authenticate request");
let mut body = authenticate_request;
let _header = RequestHeaderData::read(
&mut body,
request_header_version(ApiKey::SaslAuthenticate as i16, 2),
)
.expect("decode authenticate request header");
let request = SaslAuthenticateRequestData::read(&mut body, 2)
.expect("read authenticate request body");
server
.write_all(&sasl_authenticate_response(2))
.await
.expect("write authenticate response");
request.auth_bytes
});
task.sasl_authenticate(&mut client)
.await
.expect("sasl authenticate");
let auth_bytes = server_task.await.expect("server task");
assert_eq!(auth_bytes, Bytes::from_static(b"native-hook-payload"));
}
#[tokio::test]
async fn sasl_authenticate_uses_factory_session_hostname() {
let (mut task, client, mut server) = broker_task_with_connected_stream().await;
task.endpoint = BrokerEndpoint::from_resolved(
7,
"broker.example.com".to_owned(),
9092,
task.endpoint.addr,
);
task.config.security.protocol = SecurityProtocol::SaslPlaintext;
task.config.sasl = task
.config
.sasl
.clone()
.client_authenticator_factory(SessionPayloadFactory);
let mut client: BrokerStream = Box::new(client);
let server_task = tokio::spawn(async move {
let _handshake_request = read_frame(&mut server, None, &BufferPools::new(1))
.await
.expect("read handshake request");
server
.write_all(&sasl_handshake_response(1, "PLAIN"))
.await
.expect("write handshake response");
let authenticate_request = read_frame(&mut server, None, &BufferPools::new(1))
.await
.expect("read authenticate request");
let mut body = authenticate_request;
let _header = RequestHeaderData::read(
&mut body,
request_header_version(ApiKey::SaslAuthenticate as i16, 2),
)
.expect("decode authenticate request header");
let request = SaslAuthenticateRequestData::read(&mut body, 2)
.expect("read authenticate request body");
server
.write_all(&sasl_authenticate_response(2))
.await
.expect("write authenticate response");
request.auth_bytes
});
task.sasl_authenticate(&mut client)
.await
.expect("sasl authenticate");
let auth_bytes = server_task.await.expect("server task");
assert_eq!(auth_bytes, Bytes::from_static(b"7:broker.example.com:9092"));
}
#[tokio::test]
async fn serve_connection_closes_when_command_channel_is_closed() {
let (mut task, client, _server) = broker_task_with_connected_stream().await;
let client: BrokerStream = Box::new(client);
let mut pending = VecDeque::new();
let mut rx_open = true;
assert!(matches!(
task.serve_connection(
client,
api_versions_capabilities(),
&mut pending,
&mut rx_open
)
.await,
ServeOutcome::Closed
));
assert!(!rx_open);
}
#[tokio::test]
async fn write_command_returns_backpressure_when_pipeline_has_no_capacity() {
let (task, client, _server) = broker_task_with_connected_stream().await;
let client: BrokerStream = Box::new(client);
let (_reader, writer) = tokio::io::split(client);
let mut writer = BufWriter::new(writer);
let pipeline = Arc::new(Mutex::new(RequestPipeline::new(1, Duration::from_secs(1))));
let (reserved_tx, _reserved_rx) = oneshot::channel();
let _reserved = lock_pipeline(&pipeline)
.reserve(ApiKey::ApiVersions, 3, reserved_tx)
.expect("reserve only slot");
let (tx, rx) = oneshot::channel();
let command = request_command_with_sender(tx, Instant::now());
let wrote = task
.write_command(
&mut writer,
&pipeline,
command,
&api_versions_capabilities(),
)
.await
.expect("write command");
assert!(!wrote);
assert!(matches!(
rx.await.expect("backpressure response"),
Err(WireError::Backpressure)
));
}
#[tokio::test]
async fn read_frame_rejects_negative_and_oversized_lengths() {
let buffers = BufferPools::new(1);
let negative_bytes = (-1_i32).to_be_bytes();
let oversized_bytes = (frame::MAX_FRAME_LENGTH.saturating_add(1)).to_be_bytes();
let mut negative = &negative_bytes[..];
let mut oversized = &oversized_bytes[..];
assert!(matches!(
read_frame(&mut negative, None, &buffers).await,
Err(WireError::ConnectionClosed)
));
assert!(matches!(
read_frame(&mut oversized, None, &buffers).await,
Err(WireError::ConnectionClosed)
));
}
#[tokio::test]
async fn read_frame_reads_payload_and_releases_reusable_buffer() {
let buffers = BufferPools::new(1);
let mut framed = Vec::new();
framed.extend_from_slice(&3_i32.to_be_bytes());
framed.extend_from_slice(b"abc");
let mut reader = &framed[..];
let payload = read_frame(&mut reader, Some(16), &buffers)
.await
.expect("payload frame");
assert_eq!(payload, Bytes::from_static(b"abc"));
assert_eq!(buffers.stats().read_reused, 0);
assert_eq!(buffers.stats().read_released, 1);
}
#[tokio::test]
async fn read_response_frames_fails_inflight_and_notifies_on_disconnect() {
let (_task, client, server) = broker_task_with_connected_stream().await;
let (reader, _writer) = server.into_split();
let pipeline = Arc::new(Mutex::new(RequestPipeline::new(1, Duration::from_secs(1))));
let (response_tx, response_rx) = oneshot::channel();
let _correlation = lock_pipeline(&pipeline)
.reserve(ApiKey::ApiVersions, 3, response_tx)
.expect("reserve slot");
let disconnect = Arc::new(Notify::new());
let reader_task = tokio::spawn(read_response_frames(
reader,
Arc::clone(&pipeline),
Arc::clone(&disconnect),
Some(16),
Arc::new(BufferPools::new(1)),
));
drop(client);
reader_task.await.expect("reader task");
assert!(
response_rx
.await
.expect("in-flight completion delivered")
.is_err()
);
disconnect.notified().await;
}
}