mod handshake;
mod versions;
pub use handshake::Handshake;
pub use versions::NegotiatedVersions;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicI32, Ordering};
use bytes::Bytes;
use futures::{SinkExt, StreamExt};
use tokio::sync::{mpsc, oneshot};
use tokio_util::codec::Framed;
use tracing::debug;
use crate::error::{KafkaError, Result};
use crate::transport::{NetworkStream, TcpNetworkStream, TlsNetworkStream};
use crate::wire::{KafkaCodec, KafkaFrame};
use kafka_client_protocol::{self as protocol, Request, Response};
enum Command {
SendRequest {
data: Bytes,
correlation_id: i32,
response_tx: oneshot::Sender<Result<Bytes>>,
},
Shutdown {
done_tx: oneshot::Sender<()>,
},
}
#[derive(Clone)]
pub struct ConnectionHandle {
cmd_tx: mpsc::UnboundedSender<Command>,
negotiated: Arc<NegotiatedVersions>,
next_correlation_id: Arc<AtomicI32>,
client_id: Option<String>,
}
impl ConnectionHandle {
fn new(
cmd_tx: mpsc::UnboundedSender<Command>,
negotiated: Arc<NegotiatedVersions>,
client_id: Option<String>,
) -> Self {
Self {
cmd_tx,
negotiated,
next_correlation_id: Arc::new(AtomicI32::new(rand::random())),
client_id,
}
}
pub async fn send_request<Req, Resp>(&self, request: &Req) -> Result<Resp>
where
Req: Request,
Resp: Response,
{
let api_key = request.api_key();
let version = self
.negotiated
.get_version(api_key)
.ok_or(KafkaError::UnsupportedApi(api_key))?;
let correlation_id = self.next_correlation_id.fetch_add(1, Ordering::SeqCst);
debug!(
api_key = api_key,
version = version,
correlation_id = correlation_id,
"sending request"
);
let encoded = request.encode_frame(version, correlation_id, self.client_id.clone())?;
let (response_tx, response_rx) = oneshot::channel();
self.cmd_tx
.send(Command::SendRequest {
data: encoded,
correlation_id,
response_tx,
})
.map_err(|_| KafkaError::ConnectionClosed)?;
let response_data = response_rx
.await
.map_err(|_| KafkaError::ConnectionClosed)??;
let (header, response) = Resp::decode_frame(response_data, version)?;
if header.correlation_id() != correlation_id {
return Err(KafkaError::CorrelationIdMismatch {
expected: correlation_id,
actual: header.correlation_id(),
});
}
Ok(response)
}
pub fn negotiated(&self) -> &NegotiatedVersions {
&self.negotiated
}
pub async fn close(self) {
let (done_tx, done_rx) = oneshot::channel();
if self.cmd_tx.send(Command::Shutdown { done_tx }).is_ok() {
let _ = done_rx.await;
}
}
}
struct ConnectionReactor {
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
cmd_rx: mpsc::UnboundedReceiver<Command>,
pending: HashMap<i32, oneshot::Sender<Result<Bytes>>>,
}
impl ConnectionReactor {
fn new(
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
cmd_rx: mpsc::UnboundedReceiver<Command>,
) -> Self {
Self {
framed,
cmd_rx,
pending: HashMap::new(),
}
}
async fn run(&mut self) {
loop {
tokio::select! {
cmd = self.cmd_rx.recv() => {
match cmd {
Some(Command::SendRequest { data, correlation_id, response_tx }) => {
match self.framed.send(KafkaFrame::new(data)).await {
Ok(()) => {
if let Err(e) = self.framed.flush().await {
self.fail_pending(KafkaError::Io(e.to_string()));
return;
}
self.pending.insert(correlation_id, response_tx);
}
Err(e) => {
let _ = response_tx.send(Err(KafkaError::Io(e.to_string())));
self.fail_pending(KafkaError::ConnectionClosed);
return;
}
}
}
Some(Command::Shutdown { done_tx }) => {
debug!("ConnectionReactor received shutdown signal");
self.fail_pending(KafkaError::ConnectionClosed);
let _ = done_tx.send(());
return;
}
None => {
self.fail_pending(KafkaError::ConnectionClosed);
return;
}
}
}
frame = self.framed.next() => {
match frame {
Some(Ok(KafkaFrame { data })) => {
let corr_id = extract_correlation_id(&data);
if let Some(tx) = self.pending.remove(&corr_id) {
let _ = tx.send(Ok(data));
} else {
debug!(
"Dropping unmatched response corr_id={}, pending_count={}, data_len={}",
corr_id,
self.pending.len(),
data.len(),
);
}
}
Some(Err(e)) => {
self.fail_pending(KafkaError::Io(e.to_string()));
return;
}
None => {
self.fail_pending(KafkaError::ConnectionClosed);
return;
}
}
}
}
}
}
fn fail_pending(&mut self, err: KafkaError) {
for (_, tx) in self.pending.drain() {
let _ = tx.send(Err(err.clone()));
}
}
}
fn spawn_reactor(
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
negotiated: Arc<NegotiatedVersions>,
client_id: Option<String>,
) -> ConnectionHandle {
let (cmd_tx, cmd_rx) = mpsc::unbounded_channel();
let mut reactor = ConnectionReactor::new(framed, cmd_rx);
tokio::spawn(async move {
reactor.run().await;
});
ConnectionHandle::new(cmd_tx, negotiated, client_id)
}
pub struct SequentialConnection {
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
client_id: Option<String>,
negotiated: NegotiatedVersions,
}
impl SequentialConnection {
pub fn new(
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
client_id: Option<String>,
) -> Self {
SequentialConnection {
framed,
client_id,
negotiated: NegotiatedVersions::new(),
}
}
pub async fn send_request<Req, Resp>(&mut self, request: &Req) -> Result<Resp>
where
Req: Request,
Resp: Response,
{
let api_key = request.api_key();
let version = self.negotiated.get_version(api_key).unwrap_or(0);
let correlation_id = rand::random();
let request_data = request.encode_frame(version, correlation_id, self.client_id.clone())?;
debug!(
api_key = api_key,
version = version,
correlation_id = correlation_id,
"sending sequential request"
);
self.framed.send(KafkaFrame::new(request_data)).await?;
self.framed.flush().await?;
let frame = self
.framed
.next()
.await
.ok_or(KafkaError::ConnectionClosed)??;
debug!(
response_len = frame.data.len(),
"received sequential response"
);
let (header, response) = Resp::decode_frame(frame.data, version)?;
if header.correlation_id() != correlation_id {
return Err(KafkaError::CorrelationIdMismatch {
expected: correlation_id,
actual: header.correlation_id(),
});
}
Ok(response)
}
pub fn negotiated(&self) -> &NegotiatedVersions {
&self.negotiated
}
pub fn set_negotiated(&mut self, negotiated: NegotiatedVersions) {
self.negotiated = negotiated;
}
pub fn into_pipeline(self) -> ConnectionHandle {
spawn_reactor(self.framed, Arc::new(self.negotiated), self.client_id)
}
}
pub struct Builder {
addr: std::net::SocketAddr,
security_protocol: crate::transport::SecurityProtocol,
client_name: String,
client_version: String,
client_id: Option<String>,
sasl_config: Option<(crate::sasl::SaslMechanismType, crate::sasl::SaslCredentials)>,
}
impl Builder {
pub fn new(
addr: std::net::SocketAddr,
security_protocol: crate::transport::SecurityProtocol,
client_name: String,
client_version: String,
) -> Self {
Builder {
addr,
security_protocol,
client_name,
client_version,
client_id: Some("kafka-client".to_string()),
sasl_config: None,
}
}
pub fn with_sasl(
mut self,
mechanism: crate::sasl::SaslMechanismType,
credentials: crate::sasl::SaslCredentials,
) -> Self {
self.sasl_config = Some((mechanism, credentials));
self
}
pub fn with_client_id(mut self, client_id: String) -> Self {
self.client_id = Some(client_id);
self
}
pub async fn build(self) -> Result<ConnectionHandle> {
let tcp = tokio::net::TcpStream::connect(self.addr)
.await
.map_err(|e| KafkaError::Io(format!("Failed to connect to {}: {}", self.addr, e)))?;
let stream: Box<dyn NetworkStream> = if self.security_protocol.uses_tls() {
let tls_config = match &self.security_protocol {
crate::transport::SecurityProtocol::Ssl(cfg)
| crate::transport::SecurityProtocol::SaslSsl(cfg) => cfg.clone(),
_ => unreachable!(),
};
let tls_stream = TlsNetworkStream::from_stream(tcp, tls_config)
.await
.map_err(|e| KafkaError::TlsError {
addr: self.addr.to_string(),
details: e.to_string(),
})?;
Box::new(tls_stream) as Box<dyn NetworkStream>
} else {
Box::new(TcpNetworkStream::from_stream(tcp)) as Box<dyn NetworkStream>
};
let framed = Framed::new(stream, KafkaCodec::new());
let mut seq_conn = SequentialConnection::new(framed, self.client_id);
let negotiated =
Handshake::perform(&mut seq_conn, self.client_name, self.client_version).await?;
seq_conn.set_negotiated(negotiated);
if self.sasl_config.is_none() {
let probe_req = protocol::SaslHandshakeRequest {
mechanism: "PLAIN".to_string(),
};
match seq_conn
.send_request::<_, protocol::SaslHandshakeResponse>(&probe_req)
.await
{
Ok(resp) if !resp.mechanisms.is_empty() => {
return Err(KafkaError::AuthenticationFailed(format!(
"SASL authentication required, missing credentials. \
Server SASL mechanisms: {}. Available client: plain, scram_sha256, scram_sha512.",
resp.mechanisms.join(", ")
)));
}
Ok(_) => {
debug!("Server does not require SASL authentication");
}
Err(e) => {
return Err(KafkaError::AuthenticationFailed(format!(
"SASL authentication required, missing credentials. Probe error: {}.",
e
)));
}
}
}
if let Some((mechanism, credentials)) = self.sasl_config {
Handshake::sasl_authenticate(&mut seq_conn, mechanism, credentials)
.await
.map_err(|e| {
match &e {
KafkaError::Protocol(msg) => KafkaError::AuthenticationFailed(msg.clone()),
KafkaError::SaslError(sasl_err) => {
KafkaError::AuthenticationFailed(sasl_err.to_string())
}
_ => e,
}
})?;
}
Ok(seq_conn.into_pipeline())
}
}
fn extract_correlation_id(data: &Bytes) -> i32 {
use bytes::Buf;
(&data[..]).get_i32()
}