use std::{collections::VecDeque, fmt, sync::Arc, time::Duration};
use tokio::{
io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, BufReader},
net::TcpStream,
sync::{Mutex, mpsc, oneshot},
time,
};
use tokio_rustls::{TlsConnector, rustls::pki_types::ServerName};
use crate::native::{
KafkaClientError, KafkaClientResult,
protocol::{
API_KEY_SASL_AUTHENTICATE, API_KEY_SASL_HANDSHAKE, Encoder, FetchBodyDecoder,
SaslAuthenticateResponse, SaslHandshakeResponse, encode_sasl_authenticate_request,
encode_sasl_handshake_request, request_header_version, response_header_version,
},
security::{NativeSecurityConfig, SaslConfig, SaslMechanism, ScramClient},
};
const DEFAULT_COMMAND_BUFFER: usize = 64;
const MAX_RESPONSE_BYTES: usize = 512 * 1024 * 1024;
const TLS_RECORD_PLAINTEXT_BYTES: usize = 16 * 1024;
const TLS_ENCRYPTED_READ_BUFFER_BYTES: usize = 4 * TLS_RECORD_PLAINTEXT_BYTES;
const FETCH_STREAM_CHUNK_BYTES: usize = 64 * 1024;
#[derive(Debug, Clone)]
pub(crate) struct BrokerConnection {
sender: mpsc::Sender<Command>,
broker: String,
sasl: Option<SaslConfig>,
auth: Arc<Mutex<AuthState>>,
}
trait BrokerIo: AsyncRead + AsyncWrite + Unpin + Send {}
impl<T> BrokerIo for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
#[derive(Debug, Default)]
struct AuthState {
handshake_version: Option<i16>,
authenticate_version: Option<i16>,
authenticated: bool,
reauthenticate_at: Option<time::Instant>,
}
#[derive(Debug)]
pub(crate) struct BrokerResponse {
receiver: oneshot::Receiver<KafkaClientResult<Vec<u8>>>,
}
#[derive(Debug)]
enum Delivery {
Whole(oneshot::Sender<KafkaClientResult<Vec<u8>>>),
Decoded {
decoder: Box<FetchBodyDecoder>,
respond: oneshot::Sender<KafkaClientResult<FetchBodyDecoder>>,
},
}
impl Delivery {
fn deliver_err(self, error: KafkaClientError) {
match self {
Self::Whole(respond) => {
let _ = respond.send(Err(error));
}
Self::Decoded { respond, .. } => {
let _ = respond.send(Err(error));
}
}
}
fn deliver_empty(self) {
if let Self::Whole(respond) = self {
let _ = respond.send(Ok(Vec::new()));
}
}
}
#[derive(Debug)]
struct Command {
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
expect_response: bool,
delivery: Delivery,
}
#[derive(Debug)]
struct Pending {
correlation_id: i32,
response_header_version: i16,
delivery: Delivery,
}
enum ReadOutcome {
Whole(Vec<u8>),
Decoded,
Mismatch(i32),
}
impl BrokerConnection {
pub(crate) async fn connect(
addr: String,
broker: String,
client_id: String,
max_in_flight: usize,
request_timeout: Duration,
security: &NativeSecurityConfig,
) -> KafkaClientResult<Self> {
let stream = time::timeout(request_timeout, TcpStream::connect(&addr))
.await
.map_err(|_| {
KafkaClientError::Io(std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("Kafka TCP connect to {addr} timed out"),
))
})??;
stream.set_nodelay(true)?;
let stream: Box<dyn BrokerIo> = if let Some(tls) = &security.tls {
let server_name =
ServerName::try_from(broker.clone()).map_err(|error| KafkaClientError::Tls {
broker: broker.clone(),
message: format!("invalid TLS SNI name: {error}"),
})?;
let connector = TlsConnector::from(Arc::clone(tls));
let encrypted = BufReader::with_capacity(TLS_ENCRYPTED_READ_BUFFER_BYTES, stream);
let tls_stream =
time::timeout(request_timeout, connector.connect(server_name, encrypted))
.await
.map_err(|_| KafkaClientError::Tls {
broker: broker.clone(),
message: "TLS handshake timed out".to_owned(),
})?
.map_err(|error| KafkaClientError::Tls {
broker: broker.clone(),
message: error.to_string(),
})?;
Box::new(tls_stream)
} else {
Box::new(stream)
};
let (sender, receiver) = mpsc::channel(DEFAULT_COMMAND_BUFFER);
tokio::spawn(connection_task(
stream,
receiver,
client_id,
max_in_flight.max(1),
request_timeout,
));
Ok(Self {
sender,
broker,
sasl: security.sasl.clone(),
auth: Arc::new(Mutex::new(AuthState::default())),
})
}
pub(crate) async fn initialize_sasl(
&self,
handshake_version: i16,
authenticate_version: i16,
) -> KafkaClientResult<()> {
let Some(sasl) = self.sasl.clone() else {
return Ok(());
};
let mut auth = self.auth.lock().await;
auth.handshake_version = Some(handshake_version);
auth.authenticate_version = Some(authenticate_version);
self.authenticate(&sasl, &mut auth).await
}
pub(crate) async fn request_before_auth(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
) -> KafkaClientResult<Vec<u8>> {
self.raw_request(
api_key,
api_version,
request_header_version,
response_header_version,
body,
)
.await
}
pub(crate) async fn request(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
) -> KafkaClientResult<Vec<u8>> {
self.ensure_authenticated().await?;
self.raw_request(
api_key,
api_version,
request_header_version,
response_header_version,
body,
)
.await
}
pub(crate) async fn begin_request(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
) -> KafkaClientResult<BrokerResponse> {
self.ensure_authenticated().await?;
self.raw_begin_request(
api_key,
api_version,
request_header_version,
response_header_version,
body,
)
.await
}
async fn raw_begin_request(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
) -> KafkaClientResult<BrokerResponse> {
let receiver = self
.begin(
api_key,
api_version,
request_header_version,
response_header_version,
body,
true,
)
.await?;
Ok(BrokerResponse { receiver })
}
pub(crate) async fn send_one_way(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
body: Vec<u8>,
) -> KafkaClientResult<()> {
self.ensure_authenticated().await?;
self.begin(api_key, api_version, request_header_version, 0, body, false)
.await?
.await
.map_err(|_| KafkaClientError::ChannelClosed)?
.map(|_| ())
}
async fn raw_request(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
) -> KafkaClientResult<Vec<u8>> {
self.raw_begin_request(
api_key,
api_version,
request_header_version,
response_header_version,
body,
)
.await?
.receive()
.await
}
async fn ensure_authenticated(&self) -> KafkaClientResult<()> {
let Some(sasl) = self.sasl.clone() else {
return Ok(());
};
let mut auth = self.auth.lock().await;
if auth.authenticated
&& auth
.reauthenticate_at
.is_none_or(|deadline| deadline > time::Instant::now())
{
return Ok(());
}
self.authenticate(&sasl, &mut auth).await
}
async fn authenticate(&self, sasl: &SaslConfig, auth: &mut AuthState) -> KafkaClientResult<()> {
let handshake_version = auth.handshake_version.ok_or_else(|| {
self.sasl_error(sasl.mechanism, "SASL API versions were not negotiated")
})?;
let authenticate_version = auth.authenticate_version.ok_or_else(|| {
self.sasl_error(sasl.mechanism, "SASL API versions were not negotiated")
})?;
let body = encode_sasl_handshake_request(sasl.mechanism.name())
.map_err(|error| self.sasl_error(sasl.mechanism, error))?;
let response = self
.raw_request(
API_KEY_SASL_HANDSHAKE,
handshake_version,
request_header_version(handshake_version),
response_header_version(API_KEY_SASL_HANDSHAKE, handshake_version),
body,
)
.await
.map_err(|error| self.sasl_error(sasl.mechanism, error))?;
let handshake = SaslHandshakeResponse::decode(handshake_version, &response)
.map_err(|error| self.sasl_error(sasl.mechanism, error))?;
if handshake.error_code != 0 {
return Err(self.sasl_error(
sasl.mechanism,
format!(
"SaslHandshake returned broker code {} (advertised: {})",
handshake.error_code,
handshake.mechanisms.join(", ")
),
));
}
if !handshake
.mechanisms
.iter()
.any(|mechanism| mechanism == sasl.mechanism.name())
{
return Err(self.sasl_error(
sasl.mechanism,
format!(
"broker did not advertise the selected mechanism (advertised: {})",
handshake.mechanisms.join(", ")
),
));
}
let session_lifetime_ms = match sasl.mechanism {
SaslMechanism::Plain => {
if sasl.username.contains('\0') || sasl.password.contains('\0') {
return Err(self.sasl_error(
sasl.mechanism,
"PLAIN credentials must not contain NUL bytes",
));
}
let mut token = Vec::with_capacity(sasl.username.len() + sasl.password.len() + 2);
token.push(0);
token.extend_from_slice(sasl.username.as_bytes());
token.push(0);
token.extend_from_slice(sasl.password.as_bytes());
self.authenticate_token(sasl.mechanism, authenticate_version, &token)
.await?
.session_lifetime_ms
}
SaslMechanism::ScramSha256 | SaslMechanism::ScramSha512 => {
let client = ScramClient::new(sasl)
.map_err(|error| self.sasl_error(sasl.mechanism, error))?;
let first = self
.authenticate_token(
sasl.mechanism,
authenticate_version,
client.first_message().as_bytes(),
)
.await?;
let server_first = std::str::from_utf8(&first.auth_bytes).map_err(|_| {
self.sasl_error(sasl.mechanism, "SCRAM server-first message was not UTF-8")
})?;
let (client_final, verifier) = client
.handle_server_first(server_first)
.map_err(|error| self.sasl_error(sasl.mechanism, error))?;
let final_response = self
.authenticate_token(
sasl.mechanism,
authenticate_version,
client_final.as_bytes(),
)
.await?;
let server_final =
std::str::from_utf8(&final_response.auth_bytes).map_err(|_| {
self.sasl_error(sasl.mechanism, "SCRAM server-final message was not UTF-8")
})?;
verifier
.verify(server_final)
.map_err(|error| self.sasl_error(sasl.mechanism, error))?;
final_response.session_lifetime_ms
}
};
auth.authenticated = true;
auth.reauthenticate_at = reauthentication_deadline(session_lifetime_ms);
Ok(())
}
async fn authenticate_token(
&self,
mechanism: SaslMechanism,
version: i16,
token: &[u8],
) -> KafkaClientResult<SaslAuthenticateResponse> {
let body = encode_sasl_authenticate_request(token)
.map_err(|error| self.sasl_error(mechanism, error))?;
let response = self
.raw_request(
API_KEY_SASL_AUTHENTICATE,
version,
request_header_version(version),
response_header_version(API_KEY_SASL_AUTHENTICATE, version),
body,
)
.await
.map_err(|error| self.sasl_error(mechanism, error))?;
let response = SaslAuthenticateResponse::decode(version, &response)
.map_err(|error| self.sasl_error(mechanism, error))?;
if response.error_code != 0 {
return Err(self.sasl_error(
mechanism,
response.error_message.as_deref().map_or_else(
|| {
format!(
"SaslAuthenticate returned broker code {}",
response.error_code
)
},
|message| {
format!(
"SaslAuthenticate returned broker code {}: {message}",
response.error_code
)
},
),
));
}
Ok(response)
}
fn sasl_error(&self, mechanism: SaslMechanism, message: impl fmt::Display) -> KafkaClientError {
KafkaClientError::Sasl {
broker: self.broker.clone(),
mechanism: mechanism.name().to_owned(),
message: message.to_string(),
}
}
pub(crate) async fn request_decoded(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
decoder: FetchBodyDecoder,
) -> KafkaClientResult<FetchBodyDecoder> {
self.ensure_authenticated().await?;
let (respond, receiver) = oneshot::channel();
let command = Command {
api_key,
api_version,
request_header_version,
response_header_version,
body,
expect_response: true,
delivery: Delivery::Decoded {
decoder: Box::new(decoder),
respond,
},
};
self.sender
.send(command)
.await
.map_err(|_| KafkaClientError::ChannelClosed)?;
receiver
.await
.map_err(|_| KafkaClientError::ChannelClosed)?
}
async fn begin(
&self,
api_key: i16,
api_version: i16,
request_header_version: i16,
response_header_version: i16,
body: Vec<u8>,
expect_response: bool,
) -> KafkaClientResult<oneshot::Receiver<KafkaClientResult<Vec<u8>>>> {
let (respond, receiver) = oneshot::channel();
let command = Command {
api_key,
api_version,
request_header_version,
response_header_version,
body,
expect_response,
delivery: Delivery::Whole(respond),
};
self.sender
.send(command)
.await
.map_err(|_| KafkaClientError::ChannelClosed)?;
Ok(receiver)
}
}
fn reauthentication_deadline(session_lifetime_ms: i64) -> Option<time::Instant> {
let lifetime_ms = u64::try_from(session_lifetime_ms)
.ok()
.filter(|value| *value > 0)?;
let margin_ms = (lifetime_ms / 10).clamp(1, 30_000);
Some(time::Instant::now() + Duration::from_millis(lifetime_ms.saturating_sub(margin_ms)))
}
impl BrokerResponse {
pub(crate) async fn receive(self) -> KafkaClientResult<Vec<u8>> {
self.receiver
.await
.map_err(|_| KafkaClientError::ChannelClosed)?
}
}
async fn connection_task(
mut stream: Box<dyn BrokerIo>,
mut receiver: mpsc::Receiver<Command>,
client_id: String,
max_in_flight: usize,
request_timeout: Duration,
) {
let mut correlation_id = 1_i32;
let mut pending = VecDeque::<Pending>::new();
let mut receiving_commands = true;
let mut fetch_scratch = Vec::<u8>::new();
loop {
if !receiving_commands {
if pending.is_empty() {
break;
}
match read_response(
&mut stream,
&mut pending,
request_timeout,
&mut fetch_scratch,
)
.await
{
Ok(()) => continue,
Err(error) => {
fail_all(&mut pending, error);
break;
}
}
}
if pending.len() >= max_in_flight {
match read_response(
&mut stream,
&mut pending,
request_timeout,
&mut fetch_scratch,
)
.await
{
Ok(()) => continue,
Err(error) => {
fail_all(&mut pending, error);
break;
}
}
}
tokio::select! {
command = receiver.recv(), if receiving_commands => {
let Some(command) = command else {
receiving_commands = false;
continue;
};
let request_correlation = correlation_id;
correlation_id = if correlation_id == i32::MAX { 1 } else { correlation_id + 1 };
match write_request(&mut stream, &client_id, request_correlation, &command).await {
Ok(()) => {
if command.expect_response {
pending.push_back(Pending {
correlation_id: request_correlation,
response_header_version: command.response_header_version,
delivery: command.delivery,
});
} else {
command.delivery.deliver_empty();
}
}
Err(error) => {
command.delivery.deliver_err(error);
fail_all(&mut pending, KafkaClientError::ChannelClosed);
break;
}
}
}
result = read_response(&mut stream, &mut pending, request_timeout, &mut fetch_scratch), if !pending.is_empty() => {
if let Err(error) = result {
fail_all(&mut pending, error);
break;
}
}
}
}
}
async fn write_request<S>(
stream: &mut S,
client_id: &str,
correlation_id: i32,
command: &Command,
) -> KafkaClientResult<()>
where
S: AsyncWrite + Unpin + ?Sized,
{
let mut frame = Encoder::with_capacity(command.body.len() + 64);
frame.put_i16(command.api_key);
frame.put_i16(command.api_version);
frame.put_i32(correlation_id);
frame.put_nullable_string(Some(client_id))?;
if command.request_header_version >= 2 {
frame.put_empty_tags();
}
frame.put_raw(&command.body);
let frame = frame.into_inner();
let len = i32::try_from(frame.len())
.map_err(|_| KafkaClientError::protocol("Kafka request frame exceeds int32 length"))?;
stream.write_all(&len.to_be_bytes()).await?;
stream.write_all(&frame).await?;
Ok(())
}
async fn read_response<S>(
stream: &mut S,
pending: &mut VecDeque<Pending>,
request_timeout: Duration,
fetch_scratch: &mut Vec<u8>,
) -> KafkaClientResult<()>
where
S: AsyncRead + Unpin + ?Sized,
{
let mut len_buf = [0_u8; 4];
time::timeout(request_timeout, stream.read_exact(&mut len_buf))
.await
.map_err(|_| KafkaClientError::protocol("Kafka response read timed out"))??;
let len = i32::from_be_bytes(len_buf);
if len < 4 {
return Err(KafkaClientError::protocol(format!(
"invalid Kafka response frame length {len}"
)));
}
let len = usize::try_from(len)
.map_err(|_| KafkaClientError::protocol("Kafka response length conversion failed"))?;
if len > MAX_RESPONSE_BYTES {
return Err(KafkaClientError::protocol(format!(
"Kafka response frame too large: {len} bytes"
)));
}
let Some(front) = pending.front_mut() else {
return Err(KafkaClientError::protocol(
"Kafka response arrived with no pending request",
));
};
let response_header_version = front.response_header_version;
let expected_correlation = front.correlation_id;
let outcome = time::timeout(
request_timeout,
read_body_or_drive(
stream,
len,
response_header_version,
expected_correlation,
&mut front.delivery,
fetch_scratch,
),
)
.await
.map_err(|_| KafkaClientError::protocol("Kafka response body read timed out"))??;
let pending_request = pending
.pop_front()
.expect("pending response was checked before its body read");
match outcome {
ReadOutcome::Mismatch(correlation_id) => {
pending_request
.delivery
.deliver_err(KafkaClientError::protocol(format!(
"Kafka response correlation mismatch: expected {expected_correlation}, got {correlation_id}"
)));
Err(KafkaClientError::protocol(
"Kafka connection correlation-id mismatch",
))
}
ReadOutcome::Whole(body) => {
match pending_request.delivery {
Delivery::Whole(respond) => {
let _ = respond.send(Ok(body));
}
Delivery::Decoded { respond, .. } => {
let _ = respond.send(Err(KafkaClientError::protocol(
"Kafka Fetch response was materialized instead of streamed",
)));
}
}
Ok(())
}
ReadOutcome::Decoded => {
match pending_request.delivery {
Delivery::Decoded { decoder, respond } => {
let _ = respond.send(Ok(*decoder));
}
Delivery::Whole(respond) => {
let _ = respond.send(Err(KafkaClientError::protocol(
"Kafka response was streamed for a non-Fetch request",
)));
}
}
Ok(())
}
}
}
async fn read_body_or_drive<S>(
stream: &mut S,
frame_len: usize,
response_header_version: i16,
expected_correlation: i32,
delivery: &mut Delivery,
fetch_scratch: &mut Vec<u8>,
) -> KafkaClientResult<ReadOutcome>
where
S: AsyncRead + Unpin + ?Sized,
{
let mut remaining = frame_len;
let mut correlation = [0_u8; 4];
read_frame_bytes(stream, &mut correlation, &mut remaining).await?;
if response_header_version >= 1 {
skip_response_header_tags(stream, &mut remaining).await?;
}
let correlation_id = i32::from_be_bytes(correlation);
if correlation_id != expected_correlation {
return Ok(ReadOutcome::Mismatch(correlation_id));
}
match delivery {
Delivery::Whole(_) => {
let mut body = vec![0_u8; remaining];
read_frame_bytes(stream, &mut body, &mut remaining).await?;
Ok(ReadOutcome::Whole(body))
}
Delivery::Decoded { decoder, .. } => {
drive_fetch_decoder(stream, remaining, decoder.as_mut(), fetch_scratch).await?;
Ok(ReadOutcome::Decoded)
}
}
}
async fn drive_fetch_decoder<S>(
stream: &mut S,
mut remaining: usize,
decoder: &mut FetchBodyDecoder,
scratch: &mut Vec<u8>,
) -> KafkaClientResult<()>
where
S: AsyncRead + Unpin + ?Sized,
{
scratch.clear();
loop {
if remaining > 0 {
let want = remaining.min(FETCH_STREAM_CHUNK_BYTES);
let filled = scratch.len();
scratch.resize(filled + want, 0);
stream.read_exact(&mut scratch[filled..]).await?;
remaining -= want;
}
let more = remaining > 0;
let consumed = decoder.feed(scratch, more)?;
scratch.drain(0..consumed);
if !more {
decoder.finish()?;
return Ok(());
}
}
}
async fn skip_response_header_tags<S>(
stream: &mut S,
remaining: &mut usize,
) -> KafkaClientResult<()>
where
S: AsyncRead + Unpin + ?Sized,
{
let count = read_unsigned_varint(stream, remaining).await?;
let mut previous = None;
let mut discard = [0_u8; 8 * 1024];
for _ in 0..count {
let tag = read_unsigned_varint(stream, remaining).await?;
if previous.is_some_and(|previous| tag <= previous) {
return Err(KafkaClientError::protocol(
"Kafka response-header tagged fields were not strictly increasing",
));
}
previous = Some(tag);
let mut size = usize::try_from(read_unsigned_varint(stream, remaining).await?)
.map_err(|_| KafkaClientError::protocol("Kafka response-header tag is too large"))?;
while size > 0 {
let chunk = size.min(discard.len());
read_frame_bytes(stream, &mut discard[..chunk], remaining).await?;
size -= chunk;
}
}
Ok(())
}
async fn read_unsigned_varint<S>(stream: &mut S, remaining: &mut usize) -> KafkaClientResult<u32>
where
S: AsyncRead + Unpin + ?Sized,
{
let mut value = 0_u32;
for shift in (0..35).step_by(7) {
let mut byte = [0_u8; 1];
read_frame_bytes(stream, &mut byte, remaining).await?;
value |= u32::from(byte[0] & 0x7f) << shift;
if (byte[0] & 0x80) == 0 {
return Ok(value);
}
}
Err(KafkaClientError::protocol(
"Kafka response-header unsigned varint overflow",
))
}
async fn read_frame_bytes<S>(
stream: &mut S,
output: &mut [u8],
remaining: &mut usize,
) -> KafkaClientResult<()>
where
S: AsyncRead + Unpin + ?Sized,
{
if output.len() > *remaining {
return Err(KafkaClientError::protocol(
"Kafka response header exceeds frame length",
));
}
stream.read_exact(output).await?;
*remaining -= output.len();
Ok(())
}
fn fail_all(pending: &mut VecDeque<Pending>, error: KafkaClientError) {
let retriable = error.retriable();
let message = error.to_string();
while let Some(pending) = pending.pop_front() {
let error = if retriable {
KafkaClientError::ChannelClosed
} else {
KafkaClientError::protocol(message.clone())
};
pending.delivery.deliver_err(error);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn read_response_batches_partial_tls_records() {
let correlation_id = 41_i32;
let body = (0..(TLS_RECORD_PLAINTEXT_BYTES * 3 + 137))
.map(|index| index as u8)
.collect::<Vec<_>>();
let mut response = Vec::with_capacity(body.len() + 8);
response.extend_from_slice(
&i32::try_from(body.len() + 4)
.expect("frame length")
.to_be_bytes(),
);
response.extend_from_slice(&correlation_id.to_be_bytes());
response.extend_from_slice(&body);
let (mut reader, mut writer) = tokio::io::duplex(response.len() * 2);
let chunks = [
1,
3,
TLS_RECORD_PLAINTEXT_BYTES - 5,
TLS_RECORD_PLAINTEXT_BYTES + 11,
response.len(),
];
let write_task = tokio::spawn(async move {
let mut start = 0;
for end in chunks.into_iter().map(|end| end.min(response.len())) {
if end > start {
writer
.write_all(&response[start..end])
.await
.expect("partial response write");
tokio::task::yield_now().await;
start = end;
}
}
});
let (respond, received) = oneshot::channel();
let mut pending = VecDeque::from([Pending {
correlation_id,
response_header_version: 0,
delivery: Delivery::Whole(respond),
}]);
read_response(
&mut reader,
&mut pending,
Duration::from_secs(1),
&mut Vec::new(),
)
.await
.expect("response spanning TLS-sized records");
write_task.await.expect("response writer");
assert_eq!(
received.await.expect("response channel").expect("response"),
body
);
assert!(pending.is_empty());
}
#[tokio::test]
async fn read_response_skips_fragmented_flexible_header_tags() {
let correlation_id = 73_i32;
let body = vec![11_u8; TLS_RECORD_PLAINTEXT_BYTES + 29];
let mut payload = correlation_id.to_be_bytes().to_vec();
payload.extend_from_slice(&[2, 1, 3, 7, 8, 9, 9, 1, 10]);
payload.extend_from_slice(&body);
let mut response = i32::try_from(payload.len())
.expect("frame length")
.to_be_bytes()
.to_vec();
response.extend_from_slice(&payload);
let (mut reader, mut writer) = tokio::io::duplex(response.len() * 2);
let write_task = tokio::spawn(async move {
for chunk in response.chunks(3) {
writer.write_all(chunk).await.expect("fragmented write");
tokio::task::yield_now().await;
}
});
let (respond, received) = oneshot::channel();
let mut pending = VecDeque::from([Pending {
correlation_id,
response_header_version: 1,
delivery: Delivery::Whole(respond),
}]);
read_response(
&mut reader,
&mut pending,
Duration::from_secs(1),
&mut Vec::new(),
)
.await
.expect("flexible response");
write_task.await.expect("response writer");
assert_eq!(
received.await.expect("response channel").expect("response"),
body
);
}
}