use std::collections::HashMap;
use std::sync::atomic::{AtomicI32, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use bytes::{Buf, Bytes, BytesMut};
use futures::{SinkExt, StreamExt};
use kafka_protocol::messages::{
ApiVersionsRequest, ApiVersionsResponse, SaslAuthenticateRequest, SaslHandshakeRequest,
};
use kafka_protocol::protocol::{Decodable, StrBytes};
use tokio::net::TcpStream;
use tokio::sync::{RwLock, mpsc, oneshot, watch};
use tokio_util::codec::Framed;
use tracing::Instrument;
use crate::api_key::ApiKey;
use crate::codec;
use crate::config::ConnectionConfig;
use crate::error::{Error, Result};
use crate::error_code::ErrorCode;
use crate::rpc::Rpc;
use crate::sasl::{AuthOutcome, SaslTransport};
use crate::stats::{ConnectionStats, StatsSnapshot};
use crate::transport::Transport;
use crate::versions::{ApiVersions, our_range};
const UNSUPPORTED_VERSION: i16 = 35;
type PendingMap = HashMap<i32, oneshot::Sender<Result<Bytes>>>;
struct Pending {
waiters: Mutex<Option<PendingMap>>,
}
impl Pending {
fn new() -> Self {
Self {
waiters: Mutex::new(Some(HashMap::new())),
}
}
fn register(&self, id: i32, tx: oneshot::Sender<Result<Bytes>>, peer: &str) -> Result<()> {
let mut guard = self.waiters.lock().map_err(|_| poisoned(peer))?;
match guard.as_mut() {
Some(map) => {
map.insert(id, tx);
Ok(())
}
None => Err(Error::ConnectionClosed {
peer: peer.to_owned(),
}),
}
}
fn take(&self, id: i32) -> Option<oneshot::Sender<Result<Bytes>>> {
let mut guard = self.waiters.lock().ok()?;
guard.as_mut()?.remove(&id)
}
fn close(&self, peer: &str) {
let drained = match self.waiters.lock() {
Ok(mut guard) => guard.take(),
Err(_) => None,
};
if let Some(map) = drained {
for (_, tx) in map {
let _ = tx.send(Err(Error::ConnectionClosed {
peer: peer.to_owned(),
}));
}
}
}
fn is_closed(&self) -> bool {
self.waiters.lock().map(|g| g.is_none()).unwrap_or(true)
}
}
fn poisoned(peer: &str) -> Error {
Error::ConnectionClosed {
peer: peer.to_owned(),
}
}
#[derive(Debug, Clone)]
pub struct Connection {
inner: Arc<Inner>,
}
struct Inner {
peer: String,
node_id: Option<i32>,
config: ConnectionConfig,
versions: ApiVersions,
stats: Arc<ConnectionStats>,
commands: mpsc::Sender<BytesMut>,
pending: Arc<Pending>,
inflight: Arc<tokio::sync::Semaphore>,
correlation: AtomicI32,
reauth_gate: RwLock<()>,
shutdown: watch::Sender<bool>,
}
impl std::fmt::Debug for Inner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Connection")
.field("peer", &self.peer)
.field("node_id", &self.node_id)
.field("read_only", &self.config.read_only)
.field("stats", &self.stats.snapshot())
.finish()
}
}
impl Drop for Inner {
fn drop(&mut self) {
let _ = self.shutdown.send(true);
}
}
impl Connection {
pub async fn connect(addr: &str, config: ConnectionConfig) -> Result<Self> {
Self::connect_as(addr, None, config).await
}
pub async fn connect_as(
addr: &str,
node_id: Option<i32>,
config: ConnectionConfig,
) -> Result<Self> {
let span = tracing::debug_span!("kafka.connect", peer = addr, node_id);
async move {
let transport = tokio::time::timeout(config.connect_timeout, open(addr, &config))
.await
.map_err(|_| Error::Timeout {
api_key: ApiKey::ApiVersions,
elapsed: config.connect_timeout,
})??;
let encrypted = transport.is_encrypted();
if let Some(sasl) = &config.sasl {
sasl.check_encryption(encrypted)?;
}
let mut framed = Framed::new(transport, codec::frame_codec(config.max_frame_bytes));
let stats = ConnectionStats::new();
let mut raw = RawConn {
framed: &mut framed,
config: &config,
stats: &stats,
correlation: 0,
versions: ApiVersions::default(),
peer: addr,
};
raw.versions = negotiate_versions(&mut raw).await?;
let versions = raw.versions.clone();
let mut session_lifetime_ms = 0;
if let Some(sasl) = config.sasl.clone() {
session_lifetime_ms = crate::sasl::authenticate(&sasl, &mut raw).await?;
}
let connection = spawn(addr.to_owned(), node_id, config, versions, stats, framed);
if let Some(delay) = crate::sasl::reauth_delay(session_lifetime_ms) {
connection.spawn_reauth(delay);
}
Ok(connection)
}
.instrument(span)
.await
}
pub async fn send<R: Rpc>(&self, request: R) -> Result<R::Response> {
let deadline = Instant::now() + self.inner.config.request_timeout;
self.send_until(request, deadline).await
}
pub async fn send_until<R: Rpc>(&self, request: R, deadline: Instant) -> Result<R::Response> {
let _gate = self.inner.reauth_gate.read().await;
self.send_inner(request, deadline).await
}
async fn send_inner<R: Rpc>(&self, request: R, deadline: Instant) -> Result<R::Response> {
let api_key = R::API_KEY;
if self.inner.config.read_only && api_key.is_mutating() {
return Err(Error::ReadOnly { api_key });
}
let version = version_for::<R>(&self.inner.versions)?;
let frame = self
.round_trip(api_key, version, &request, deadline)
.await?;
codec::decode_response::<R>(api_key, version, frame)
}
async fn round_trip<R: Rpc>(
&self,
api_key: ApiKey,
version: i16,
request: &R,
deadline: Instant,
) -> Result<Bytes> {
let inner = &self.inner;
let correlation_id = inner.correlation.fetch_add(1, Ordering::Relaxed) & i32::MAX;
let span = tracing::debug_span!(
"kafka.rpc",
peer = %inner.peer,
node_id = inner.node_id,
api = %api_key,
version,
correlation_id,
);
let _entered = span.enter();
let bytes = codec::encode_request(
api_key,
version,
correlation_id,
inner.config.client_id.as_deref(),
request,
)?;
let permit = with_deadline(deadline, api_key, inner.inflight.clone().acquire_owned())
.await?
.map_err(|_| Error::ConnectionClosed {
peer: inner.peer.clone(),
})?;
let (tx, rx) = oneshot::channel();
inner.pending.register(correlation_id, tx, &inner.peer)?;
if inner.commands.send(bytes).await.is_err() {
let _ = inner.pending.take(correlation_id);
return Err(Error::ConnectionClosed {
peer: inner.peer.clone(),
});
}
let started = Instant::now();
let received = with_deadline(deadline, api_key, rx).await;
drop(permit);
metrics::histogram!("kafka_rpc_duration_seconds", "api" => api_key.name())
.record(started.elapsed().as_secs_f64());
metrics::counter!("kafka_rpc_total", "api" => api_key.name()).increment(1);
match received {
Ok(Ok(result)) => result,
Ok(Err(_recv_error)) => Err(Error::ConnectionClosed {
peer: inner.peer.clone(),
}),
Err(timeout) => {
let _ = inner.pending.take(correlation_id);
Err(timeout)
}
}
}
pub fn versions(&self) -> &ApiVersions {
&self.inner.versions
}
pub fn negotiated_for<R: Rpc>(&self) -> Result<i16> {
version_for::<R>(&self.inner.versions)
}
pub fn negotiated_version(&self, api_key: ApiKey) -> Option<i16> {
self.inner
.versions
.get(api_key)
.and_then(|e| e.negotiated())
}
pub fn stats(&self) -> &Arc<ConnectionStats> {
&self.inner.stats
}
pub fn stats_snapshot(&self) -> StatsSnapshot {
self.inner.stats.snapshot()
}
pub fn peer(&self) -> &str {
&self.inner.peer
}
pub fn node_id(&self) -> Option<i32> {
self.inner.node_id
}
pub fn is_read_only(&self) -> bool {
self.inner.config.read_only
}
pub fn is_closed(&self) -> bool {
self.inner.pending.is_closed() || self.inner.commands.is_closed()
}
pub fn close(&self) {
let _ = self.inner.shutdown.send(true);
self.inner.pending.close(&self.inner.peer);
}
fn spawn_reauth(&self, first_delay: Duration) {
let connection = self.clone();
let mut shutdown = self.inner.shutdown.subscribe();
tokio::spawn(async move {
let mut delay = first_delay;
loop {
tokio::select! {
_ = shutdown.changed() => return,
_ = tokio::time::sleep(delay) => {}
}
let Some(sasl) = connection.inner.config.sasl.clone() else {
return;
};
let _exclusive = connection.inner.reauth_gate.write().await;
let mut transport = ReauthTransport {
connection: &connection,
};
match crate::sasl::authenticate(&sasl, &mut transport).await {
Ok(lifetime) => match crate::sasl::reauth_delay(lifetime) {
Some(next) => {
tracing::debug!(peer = %connection.inner.peer, ?next, "re-authenticated");
delay = next;
}
None => return,
},
Err(error) => {
tracing::warn!(
peer = %connection.inner.peer,
%error,
"re-authentication failed, closing connection"
);
connection.close();
return;
}
}
}
});
}
}
async fn open(addr: &str, config: &ConnectionConfig) -> Result<Transport> {
let tcp = TcpStream::connect(addr)
.await
.map_err(|e| Error::transport("connecting to broker", e))?;
tcp.set_nodelay(true)
.map_err(|e| Error::transport("disabling Nagle", e))?;
match &config.tls {
None => Ok(Transport::Plain(tcp)),
Some(tls) => {
let host = addr.rsplit_once(':').map(|(h, _)| h).unwrap_or(addr);
let connector = tls.connector()?;
let server_name = tls.server_name(host)?;
let stream = connector
.connect(server_name, tcp)
.await
.map_err(|e| Error::transport("TLS handshake", e))?;
Ok(Transport::Tls(Box::new(stream)))
}
}
}
fn spawn(
peer: String,
node_id: Option<i32>,
config: ConnectionConfig,
versions: ApiVersions,
stats: Arc<ConnectionStats>,
framed: Framed<Transport, tokio_util::codec::LengthDelimitedCodec>,
) -> Connection {
let (commands_tx, mut commands_rx) = mpsc::channel::<BytesMut>(config.max_in_flight);
let (shutdown_tx, _) = watch::channel(false);
let pending = Arc::new(Pending::new());
let inflight = Arc::new(tokio::sync::Semaphore::new(config.max_in_flight));
let (mut sink, mut stream) = framed.split();
{
let stats = stats.clone();
let pending = pending.clone();
let peer = peer.clone();
let mut shutdown = shutdown_tx.subscribe();
tokio::spawn(async move {
loop {
let frame = tokio::select! {
_ = shutdown.changed() => break,
frame = commands_rx.recv() => match frame {
Some(frame) => frame,
None => break,
},
};
let len = frame.len() + 4;
if let Err(error) = sink.send(frame.freeze()).await {
tracing::debug!(%peer, %error, "write failed");
break;
}
stats.record_sent(len);
}
pending.close(&peer);
});
}
{
let stats = stats.clone();
let pending = pending.clone();
let peer = peer.clone();
let mut shutdown = shutdown_tx.subscribe();
tokio::spawn(async move {
loop {
let frame = tokio::select! {
_ = shutdown.changed() => break,
frame = stream.next() => match frame {
Some(frame) => frame,
None => break,
},
};
match frame {
Ok(frame) => {
stats.record_received(frame.len() + 4);
let frame = frame.freeze();
match codec::peek_correlation_id(&frame) {
Ok(correlation_id) => match pending.take(correlation_id) {
Some(waiter) => {
let _ = waiter.send(Ok(frame));
}
None => tracing::trace!(
%peer,
correlation_id,
"response for a caller that went away"
),
},
Err(error) => {
tracing::warn!(%peer, %error, "undecodable response frame");
break;
}
}
}
Err(error) => {
tracing::debug!(%peer, %error, "read failed");
break;
}
}
}
pending.close(&peer);
});
}
Connection {
inner: Arc::new(Inner {
peer,
node_id,
config,
versions,
stats,
commands: commands_tx,
pending,
inflight,
correlation: AtomicI32::new(1),
reauth_gate: RwLock::new(()),
shutdown: shutdown_tx,
}),
}
}
fn typed_range<R: Rpc>() -> crate::versions::VersionRange {
let request = R::VERSIONS;
let response = <R::Response as kafka_protocol::protocol::Message>::VERSIONS;
crate::versions::VersionRange::new(request.min.max(response.min), request.max.min(response.max))
}
fn version_for<R: Rpc>(versions: &ApiVersions) -> Result<i16> {
versions.negotiate_with(R::API_KEY, Some(typed_range::<R>()))
}
async fn with_deadline<F: Future>(
deadline: Instant,
api_key: ApiKey,
future: F,
) -> Result<F::Output> {
let started = Instant::now();
match tokio::time::timeout_at(deadline.into(), future).await {
Ok(value) => Ok(value),
Err(_) => Err(Error::Timeout {
api_key,
elapsed: started.elapsed(),
}),
}
}
struct RawConn<'a> {
framed: &'a mut Framed<Transport, tokio_util::codec::LengthDelimitedCodec>,
config: &'a ConnectionConfig,
stats: &'a Arc<ConnectionStats>,
correlation: i32,
versions: ApiVersions,
peer: &'a str,
}
impl RawConn<'_> {
async fn round_trip<R: Rpc>(
&mut self,
api_key: ApiKey,
version: i16,
request: &R,
) -> Result<Bytes> {
self.correlation = self.correlation.wrapping_add(1) & i32::MAX;
let bytes = codec::encode_request(
api_key,
version,
self.correlation,
self.config.client_id.as_deref(),
request,
)?;
let len = bytes.len() + 4;
let deadline = Instant::now() + self.config.connect_timeout;
with_deadline(deadline, api_key, self.framed.send(bytes.freeze()))
.await?
.map_err(|e| Error::transport("sending handshake request", e))?;
self.stats.record_sent(len);
let frame = with_deadline(deadline, api_key, self.framed.next())
.await?
.ok_or_else(|| Error::ConnectionClosed {
peer: self.peer.to_owned(),
})?
.map_err(|e| Error::transport("reading handshake response", e))?;
self.stats.record_received(frame.len() + 4);
let frame = frame.freeze();
let correlation = codec::peek_correlation_id(&frame)?;
if correlation != self.correlation {
return Err(Error::decode(
"handshake response correlation id",
std::io::Error::other(format!("expected {}, got {correlation}", self.correlation)),
));
}
Ok(frame)
}
async fn call<R: Rpc>(&mut self, api_key: ApiKey, request: &R) -> Result<R::Response> {
let version = version_for::<R>(&self.versions)?;
let frame = self.round_trip(api_key, version, request).await?;
codec::decode_response::<R>(api_key, version, frame)
}
}
async fn negotiate_versions(raw: &mut RawConn<'_>) -> Result<ApiVersions> {
let our_max = our_range(ApiKey::ApiVersions)
.map(|r| r.max)
.ok_or_else(|| Error::Unsupported("this build cannot encode ApiVersions".to_owned()))?;
let request = ApiVersionsRequest::default()
.with_client_software_name(StrBytes::from_string(
raw.config.client_software_name.clone(),
))
.with_client_software_version(StrBytes::from_string(
raw.config.client_software_version.clone(),
));
let frame = raw
.round_trip(ApiKey::ApiVersions, our_max, &request)
.await?;
let body = codec::split_response_body(ApiKey::ApiVersions, our_max, frame)?;
let error_code = peek_error_code(&body)?;
let response = if error_code == UNSUPPORTED_VERSION {
tracing::debug!(
peer = raw.peer,
our_max,
"broker rejected ApiVersions at our maximum; falling back to v0"
);
let mut at_v0 = body.clone();
let downgraded = ApiVersionsResponse::decode(&mut at_v0, 0)
.map_err(|e| Error::decode("decoding ApiVersions v0 fallback", e))?;
if downgraded.api_keys.is_empty() {
let frame = raw.round_trip(ApiKey::ApiVersions, 0, &request).await?;
let mut body = codec::split_response_body(ApiKey::ApiVersions, 0, frame)?;
ApiVersionsResponse::decode(&mut body, 0)
.map_err(|e| Error::decode("decoding ApiVersions v0", e))?
} else {
downgraded
}
} else {
let mut body = body;
ApiVersionsResponse::decode(&mut body, our_max)
.map_err(|e| Error::decode("decoding ApiVersions", e))?
};
if let Some(code) = ErrorCode::from_code(response.error_code) {
return Err(Error::from_code(code, None));
}
Ok(ApiVersions::from_triples(
response
.api_keys
.iter()
.map(|k| (k.api_key, k.min_version, k.max_version)),
))
}
fn peek_error_code(body: &Bytes) -> Result<i16> {
if body.len() < 2 {
return Err(Error::decode(
"reading ApiVersions error code",
std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"ApiVersions response body is shorter than its error code",
),
));
}
let mut head = body.get(..2).unwrap_or_default();
Ok(head.get_i16())
}
impl SaslTransport for RawConn<'_> {
async fn handshake(&mut self, mechanism: &str) -> Result<Vec<String>> {
let version = version_for::<SaslHandshakeRequest>(&self.versions)?;
if version < 1 {
return Err(Error::Unsupported(
"broker only offers SaslHandshake v0, which uses raw token framing".to_owned(),
));
}
let request = SaslHandshakeRequest::default()
.with_mechanism(StrBytes::from_string(mechanism.to_owned()));
let frame = self
.round_trip(ApiKey::SaslHandshake, version, &request)
.await?;
let response =
codec::decode_response::<SaslHandshakeRequest>(ApiKey::SaslHandshake, version, frame)?;
let mechanisms: Vec<String> = response
.mechanisms
.iter()
.map(|m| m.as_str().to_owned())
.collect();
if let Some(code) = ErrorCode::from_code(response.error_code) {
return Err(Error::Authentication(format!(
"broker rejected mechanism {mechanism} ({code}); it enables [{}]",
mechanisms.join(", ")
)));
}
Ok(mechanisms)
}
async fn authenticate(&mut self, token: Vec<u8>) -> Result<AuthOutcome> {
let request = SaslAuthenticateRequest::default().with_auth_bytes(Bytes::from(token));
let response = self.call(ApiKey::SaslAuthenticate, &request).await?;
auth_outcome(
response.error_code,
response.error_message.map(|m| m.as_str().to_owned()),
response.auth_bytes,
response.session_lifetime_ms,
)
}
}
struct ReauthTransport<'a> {
connection: &'a Connection,
}
impl SaslTransport for ReauthTransport<'_> {
async fn handshake(&mut self, mechanism: &str) -> Result<Vec<String>> {
let request = SaslHandshakeRequest::default()
.with_mechanism(StrBytes::from_string(mechanism.to_owned()));
let deadline = Instant::now() + self.connection.inner.config.request_timeout;
let response = self.connection.send_inner(request, deadline).await?;
let mechanisms: Vec<String> = response
.mechanisms
.iter()
.map(|m| m.as_str().to_owned())
.collect();
if let Some(code) = ErrorCode::from_code(response.error_code) {
return Err(Error::Authentication(format!(
"broker rejected mechanism {mechanism} on re-authentication ({code})"
)));
}
Ok(mechanisms)
}
async fn authenticate(&mut self, token: Vec<u8>) -> Result<AuthOutcome> {
let request = SaslAuthenticateRequest::default().with_auth_bytes(Bytes::from(token));
let deadline = Instant::now() + self.connection.inner.config.request_timeout;
let response = self.connection.send_inner(request, deadline).await?;
auth_outcome(
response.error_code,
response.error_message.map(|m| m.as_str().to_owned()),
response.auth_bytes,
response.session_lifetime_ms,
)
}
}
fn auth_outcome(
error_code: i16,
error_message: Option<String>,
auth_bytes: Bytes,
session_lifetime_ms: i64,
) -> Result<AuthOutcome> {
if let Some(code) = ErrorCode::from_code(error_code) {
return Err(Error::from_code(code, error_message));
}
Ok(AuthOutcome {
auth_bytes: auth_bytes.to_vec(),
session_lifetime_ms,
})
}
#[cfg(test)]
mod tests {
use super::*;
use kafka_protocol::messages::ResponseHeader;
use kafka_protocol::protocol::Encodable;
#[test]
fn negotiation_clamps_to_the_request_type_not_the_api_key() {
use kafka_protocol::messages::OffsetFetchRequest;
let by_api_key = crate::versions::our_range(ApiKey::OffsetFetch).expect("known key");
let typed = typed_range::<OffsetFetchRequest>();
assert_eq!(by_api_key.max, 10, "ApiKey::valid_versions reports v10");
assert_eq!(typed.max, 9, "but the request encoder stops at v9");
let table = ApiVersions::from_triples([(ApiKey::OffsetFetch.code(), 1, 10)]);
assert_eq!(
table.negotiate_with(ApiKey::OffsetFetch, Some(typed)).ok(),
Some(9)
);
assert_eq!(version_for::<OffsetFetchRequest>(&table).ok(), Some(9));
assert_eq!(table.negotiate(ApiKey::OffsetFetch).ok(), Some(10));
}
#[test]
fn a_typed_range_is_the_overlap_of_request_and_response() {
use kafka_protocol::messages::MetadataRequest;
let typed = typed_range::<MetadataRequest>();
let by_api_key = crate::versions::our_range(ApiKey::Metadata).expect("known key");
assert_eq!((typed.min, typed.max), (by_api_key.min, by_api_key.max));
}
#[test]
fn every_request_negotiates_to_a_version_its_own_encoder_accepts() {
macro_rules! check {
($($ty:ty),+ $(,)?) => {$({
let key = <$ty as Rpc>::API_KEY;
let ours = crate::versions::our_range(key).expect("this build knows the key");
let table = ApiVersions::from_triples([(key.code(), ours.min, ours.max)]);
let picked = match version_for::<$ty>(&table) {
Ok(version) => version,
Err(error) => panic!("{key}: no usable version: {error}"),
};
let request = <$ty as kafka_protocol::protocol::Message>::VERSIONS;
assert!(
picked >= request.min && picked <= request.max,
"{key}: send path picked v{picked}, but the request encoder only \
covers v{}..=v{}",
request.min,
request.max,
);
})+};
}
use kafka_protocol::messages::*;
check!(
ApiVersionsRequest,
MetadataRequest,
FetchRequest,
ListOffsetsRequest,
FindCoordinatorRequest,
OffsetFetchRequest,
OffsetCommitRequest,
OffsetDeleteRequest,
DescribeGroupsRequest,
ListGroupsRequest,
DeleteGroupsRequest,
ConsumerGroupDescribeRequest,
ShareGroupDescribeRequest,
CreateTopicsRequest,
DeleteTopicsRequest,
CreatePartitionsRequest,
DeleteRecordsRequest,
DescribeTopicPartitionsRequest,
DescribeConfigsRequest,
IncrementalAlterConfigsRequest,
DescribeClusterRequest,
DescribeLogDirsRequest,
DescribeAclsRequest,
CreateAclsRequest,
DeleteAclsRequest,
DescribeClientQuotasRequest,
AlterClientQuotasRequest,
DescribeUserScramCredentialsRequest,
AlterUserScramCredentialsRequest,
ListPartitionReassignmentsRequest,
AlterPartitionReassignmentsRequest,
ElectLeadersRequest,
ListTransactionsRequest,
DescribeTransactionsRequest,
DescribeProducersRequest,
SaslHandshakeRequest,
SaslAuthenticateRequest,
);
}
#[test]
fn pending_resolves_every_waiter_when_the_socket_dies() {
let pending = Pending::new();
let (tx, mut rx) = oneshot::channel();
pending.register(7, tx, "broker:9092").unwrap();
pending.close("broker:9092");
let received = rx.try_recv().expect("resolved, not dropped");
assert!(matches!(received, Err(Error::ConnectionClosed { .. })));
assert!(pending.is_closed());
}
#[test]
fn registering_on_a_dead_connection_fails_immediately() {
let pending = Pending::new();
pending.close("broker:9092");
let (tx, _rx) = oneshot::channel();
let err = pending.register(1, tx, "broker:9092").unwrap_err();
assert!(matches!(err, Error::ConnectionClosed { .. }));
}
#[test]
fn taking_an_unknown_correlation_id_is_not_an_error() {
let pending = Pending::new();
assert!(pending.take(99).is_none());
assert!(!pending.is_closed());
}
fn api_versions_body(version: i16, error_code: i16) -> BytesMut {
let mut body = BytesMut::new();
body.extend_from_slice(&error_code.to_be_bytes());
if version >= 3 {
body.extend_from_slice(&[1]); body.extend_from_slice(&0i32.to_be_bytes()); body.extend_from_slice(&[0]); } else {
body.extend_from_slice(&0i32.to_be_bytes()); if version >= 1 {
body.extend_from_slice(&0i32.to_be_bytes()); }
}
body
}
#[test]
fn an_api_versions_error_body_still_yields_a_readable_error_code() {
let body = api_versions_body(0, UNSUPPORTED_VERSION);
assert_eq!(
peek_error_code(&body.freeze()).unwrap(),
UNSUPPORTED_VERSION
);
}
#[test]
fn a_truncated_api_versions_body_is_an_error_not_a_panic() {
assert!(peek_error_code(&Bytes::from_static(&[0])).is_err());
}
#[test]
fn the_error_code_is_readable_before_the_body_version_is_known() {
for version in [0i16, 3] {
let mut frame = BytesMut::new();
ResponseHeader::default()
.with_correlation_id(1)
.encode(&mut frame, 0)
.unwrap();
frame.extend_from_slice(&api_versions_body(version, UNSUPPORTED_VERSION));
let body =
codec::split_response_body(ApiKey::ApiVersions, version, frame.freeze()).unwrap();
assert_eq!(peek_error_code(&body).unwrap(), UNSUPPORTED_VERSION);
}
}
#[test]
fn a_v0_fallback_body_decodes_into_a_version_table() {
let mut body = BytesMut::new();
body.extend_from_slice(&UNSUPPORTED_VERSION.to_be_bytes());
body.extend_from_slice(&1i32.to_be_bytes()); body.extend_from_slice(&ApiKey::Metadata.code().to_be_bytes());
body.extend_from_slice(&0i16.to_be_bytes());
body.extend_from_slice(&13i16.to_be_bytes());
let mut buf = body.freeze();
let decoded = ApiVersionsResponse::decode(&mut buf, 0).unwrap();
assert_eq!(decoded.error_code, UNSUPPORTED_VERSION);
assert_eq!(decoded.api_keys.len(), 1, "the table survives the error");
let table = ApiVersions::from_triples(
decoded
.api_keys
.iter()
.map(|k| (k.api_key, k.min_version, k.max_version)),
);
assert!(table.supports(ApiKey::Metadata));
}
}