mod handshake;
mod versions;
pub use handshake::Handshake;
pub use versions::NegotiatedVersions;
use std::sync::Arc;
use std::sync::atomic::{AtomicI32, Ordering};
use std::time::Duration;
use futures::{SinkExt, StreamExt};
use tokio_util::codec::Framed;
use tracing::debug;
use crate::error::{KafkaError, Result};
use crate::transport::NetworkStream;
use crate::wire::{KafkaCodec, KafkaFrame};
use kafka_client_protocol::{Request, Response};
pub struct Connection {
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
negotiated: Arc<NegotiatedVersions>,
next_correlation_id: AtomicI32,
}
impl Connection {
pub(crate) fn new(
framed: Framed<Box<dyn NetworkStream>, KafkaCodec>,
negotiated: Arc<NegotiatedVersions>,
) -> Self {
Connection {
framed,
negotiated,
next_correlation_id: AtomicI32::new(rand::random()),
}
}
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)
.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 request_data =
request.encode_frame(version, correlation_id, Some("kafka-client".to_string()))?;
self.framed.send(KafkaFrame::new(request_data)).await?;
self.framed.flush().await?;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
let frame = tokio::time::timeout(REQUEST_TIMEOUT, self.framed.next())
.await
.map_err(|_| {
debug!(api_key, "request timed out after 30s");
KafkaError::RequestTimeout
})?
.ok_or(KafkaError::ConnectionClosed)??;
debug!(response_len = frame.data.len(), "received 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 async fn send_request_at<Req, Resp>(&mut self, request: &Req, version: i16) -> Result<Resp>
where
Req: Request,
Resp: Response,
{
let api_key = request.api_key();
let correlation_id = self.next_correlation_id.fetch_add(1, Ordering::SeqCst);
debug!(
api_key = api_key,
version = version,
"sending request at explicit version"
);
let request_data =
request.encode_frame(version, correlation_id, Some("kafka-client".to_string()))?;
self.framed.send(KafkaFrame::new(request_data)).await?;
self.framed.flush().await?;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
let frame = tokio::time::timeout(REQUEST_TIMEOUT, self.framed.next())
.await
.map_err(|_| {
debug!(api_key, version, "request timed out after 30s");
KafkaError::RequestTimeout
})?
.ok_or(KafkaError::ConnectionClosed)??;
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 async fn close(mut self) -> Result<()> {
if let Err(e) = self.framed.close().await {
debug!("Error closing framed connection: {}", e);
}
Ok(())
}
}
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<Connection> {
use crate::transport::TransportConnector;
let stream = TransportConnector::connect(self.addr, &self.security_protocol).await?;
let framed = Framed::new(stream, crate::wire::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 let Some((mechanism, credentials)) = self.sasl_config {
Handshake::sasl_authenticate(&mut seq_conn, mechanism, credentials).await?;
}
Ok(seq_conn.into_pipeline())
}
}
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) -> Connection {
Connection::new(self.framed, Arc::new(self.negotiated))
}
pub fn into_parts(self) -> Framed<Box<dyn NetworkStream>, KafkaCodec> {
self.framed
}
}