#![cfg_attr(not(feature = "rdkafka"), allow(dead_code))]
use std::sync::{Arc, Mutex};
#[cfg(feature = "rdkafka")]
use std::{
sync::atomic::{AtomicU64, Ordering},
time::{Duration, Instant},
};
use datum::{Flow, Keep, NotUsed, Sink, StreamCompletion, StreamError};
#[cfg(feature = "rdkafka")]
use rdkafka::{
ClientConfig, ClientContext,
error::{KafkaError, RDKafkaErrorCode},
message::{Header, Message, OwnedHeaders},
producer::{BaseProducer, BaseRecord, DeliveryResult, Producer, ProducerContext},
util::Timeout,
};
use crate::{
KafkaMetrics, KafkaProducerBackend, KafkaProducerSettings, MqError, MqResult, ProducerRecord,
native::{NativeKafkaProducerControl, NativeKafkaProducerHandle},
};
const PRODUCER_BATCH_SIZE: usize = 256;
#[cfg(feature = "rdkafka")]
const PRODUCER_POLL_SLICE: Duration = Duration::from_millis(1);
pub struct KafkaSink;
impl KafkaSink {
#[must_use]
pub fn plain(settings: KafkaProducerSettings) -> Sink<ProducerRecord, KafkaProducerControl> {
Flow::identity()
.grouped(PRODUCER_BATCH_SIZE)
.to_mat(Self::batched(settings), Keep::right)
}
fn batched(settings: KafkaProducerSettings) -> Sink<Vec<ProducerRecord>, KafkaProducerControl> {
match settings.producer_backend {
#[cfg(feature = "rdkafka")]
KafkaProducerBackend::Rdkafka => Self::rdkafka_batched(settings),
KafkaProducerBackend::Native => Self::native_batched(settings),
}
}
#[cfg(feature = "rdkafka")]
fn rdkafka_batched(
settings: KafkaProducerSettings,
) -> Sink<Vec<ProducerRecord>, KafkaProducerControl> {
Sink::setup(move |_materializer, _attributes| {
let control = KafkaProducerControl::new_rdkafka(settings.clone())
.expect("Kafka producer configuration must be valid at materialization");
let runner = control
.rdkafka_runner()
.expect("rdkafka producer control owns an rdkafka runner");
Sink::foreach_result(move |records| {
runner
.lock()
.map_err(|_| StreamError::Failed("Kafka producer lock poisoned".to_owned()))?
.send_batch(records)
.map_err(StreamError::from)
})
.map_materialized_value(move |completion| {
control.attach_completion(completion);
control.clone()
})
})
}
fn native_batched(
settings: KafkaProducerSettings,
) -> Sink<Vec<ProducerRecord>, KafkaProducerControl> {
Sink::setup(move |_materializer, _attributes| {
let metrics = KafkaMetrics::default();
let (native, handle) = NativeKafkaProducerControl::start(&settings, metrics.clone())
.expect("native Kafka producer configuration must be valid at materialization");
let control = KafkaProducerControl::new_native(native, metrics);
Sink::foreach_result(move |records| {
let handle: NativeKafkaProducerHandle = handle.clone();
handle.send(records).map_err(StreamError::from)
})
.map_materialized_value(move |completion| {
control.attach_completion(completion);
control.clone()
})
})
}
}
#[derive(Clone)]
pub struct KafkaProducerControl {
state: Arc<ProducerControlState>,
}
struct ProducerControlState {
backend: ProducerControlBackend,
metrics: KafkaMetrics,
completion: Mutex<Option<StreamCompletion<NotUsed>>>,
}
enum ProducerControlBackend {
#[cfg(feature = "rdkafka")]
Rdkafka(Arc<Mutex<ProducerRunner>>),
Native(NativeKafkaProducerControl),
}
impl KafkaProducerControl {
#[cfg(feature = "rdkafka")]
fn new_rdkafka(settings: KafkaProducerSettings) -> MqResult<Self> {
let metrics = KafkaMetrics::default();
let delivery = Arc::new(DeliveryState::new(metrics.clone()));
let context = DeliveryContext {
state: Arc::clone(&delivery),
};
let producer: BaseProducer<DeliveryContext> =
settings.to_client_config().create_with_context(context)?;
Ok(Self {
state: Arc::new(ProducerControlState {
backend: ProducerControlBackend::Rdkafka(Arc::new(Mutex::new(ProducerRunner {
producer,
delivery,
metrics: metrics.clone(),
settings,
enqueued_since_poll: 0,
}))),
metrics,
completion: Mutex::new(None),
}),
})
}
fn new_native(native: NativeKafkaProducerControl, metrics: KafkaMetrics) -> Self {
Self {
state: Arc::new(ProducerControlState {
backend: ProducerControlBackend::Native(native),
metrics,
completion: Mutex::new(None),
}),
}
}
#[cfg(feature = "rdkafka")]
fn rdkafka_runner(&self) -> Option<Arc<Mutex<ProducerRunner>>> {
match &self.state.backend {
ProducerControlBackend::Rdkafka(runner) => Some(Arc::clone(runner)),
ProducerControlBackend::Native(_) => None,
}
}
fn attach_completion(&self, completion: StreamCompletion<NotUsed>) {
if let Ok(mut slot) = self.state.completion.lock() {
*slot = Some(completion);
}
}
#[must_use]
pub fn metrics(&self) -> KafkaMetrics {
self.state.metrics.clone()
}
pub fn drain_and_shutdown(&self) -> MqResult<()> {
let completion = self
.state
.completion
.lock()
.map_err(|_| MqError::Failed("Kafka producer completion lock poisoned".to_owned()))?
.take();
let completion_result = completion.map_or(Ok(()), |completion| {
completion
.wait()
.map(|_| ())
.map_err(|error| MqError::Failed(error.to_string()))
});
let finish_result = match &self.state.backend {
#[cfg(feature = "rdkafka")]
ProducerControlBackend::Rdkafka(runner) => runner
.lock()
.map_err(|_| MqError::Failed("Kafka producer lock poisoned".to_owned()))?
.finish(),
ProducerControlBackend::Native(native) => native.finish(),
};
completion_result.and(finish_result)
}
pub fn flush(&self) -> MqResult<()> {
match &self.state.backend {
#[cfg(feature = "rdkafka")]
ProducerControlBackend::Rdkafka(runner) => runner
.lock()
.map_err(|_| MqError::Failed("Kafka producer lock poisoned".to_owned()))?
.finish(),
ProducerControlBackend::Native(native) => native.flush(),
}
}
}
#[cfg(feature = "rdkafka")]
struct ProducerRunner {
producer: BaseProducer<DeliveryContext>,
delivery: Arc<DeliveryState>,
metrics: KafkaMetrics,
settings: KafkaProducerSettings,
enqueued_since_poll: usize,
}
#[cfg(feature = "rdkafka")]
impl ProducerRunner {
fn send_batch(&mut self, records: Vec<ProducerRecord>) -> MqResult<()> {
for record in records {
self.send_record(&record)?;
}
self.poll(Duration::ZERO)
}
fn send_record(&mut self, record: &ProducerRecord) -> MqResult<()> {
self.wait_for_capacity()?;
let mut base_record = to_base_record(record);
let deadline = Instant::now() + self.settings.queue_timeout;
loop {
match self.producer.send(base_record) {
Ok(()) => {
let in_flight = self.delivery.sent();
self.metrics.producer_in_flight(in_flight);
self.enqueued_since_poll += 1;
if self.enqueued_since_poll >= PRODUCER_BATCH_SIZE {
self.poll(Duration::ZERO)?;
}
return Ok(());
}
Err((error, returned)) if is_queue_full(&error) && Instant::now() < deadline => {
self.metrics.queue_full();
base_record = returned;
self.poll(PRODUCER_POLL_SLICE)?;
}
Err((error, _returned)) => {
self.metrics.delivery_failed();
return Err(MqError::Kafka(error));
}
}
}
}
fn wait_for_capacity(&mut self) -> MqResult<()> {
if self.delivery.in_flight() < self.settings.in_flight_limit as u64 {
return Ok(());
}
let deadline = Instant::now() + self.settings.queue_timeout;
while self.delivery.in_flight() >= self.settings.in_flight_limit as u64
|| self.producer.in_flight_count().max(0) as usize >= self.settings.in_flight_limit
{
if Instant::now() >= deadline {
return Err(MqError::Failed(format!(
"Kafka producer capacity wait timed out with {} records in flight",
self.delivery.in_flight()
)));
}
self.poll(PRODUCER_POLL_SLICE)?;
}
Ok(())
}
fn poll(&mut self, timeout: Duration) -> MqResult<()> {
self.producer.poll(Timeout::After(timeout));
self.enqueued_since_poll = 0;
self.check_delivery_error()
}
fn finish(&mut self) -> MqResult<()> {
let deadline = Instant::now() + self.settings.drain_timeout;
while self.delivery.in_flight() > 0 {
if Instant::now() >= deadline {
return Err(MqError::DrainTimeout);
}
self.poll(Duration::from_millis(10))?;
}
self.producer
.flush(Timeout::After(self.settings.flush_timeout))?;
self.poll(Duration::ZERO)
}
fn check_delivery_error(&self) -> MqResult<()> {
self.delivery
.take_error()
.map_or(Ok(()), |error| Err(MqError::Delivery(error)))
}
}
#[derive(Clone)]
#[cfg(feature = "rdkafka")]
struct DeliveryContext {
state: Arc<DeliveryState>,
}
#[cfg(feature = "rdkafka")]
impl ClientContext for DeliveryContext {}
#[cfg(feature = "rdkafka")]
impl ProducerContext for DeliveryContext {
type DeliveryOpaque = ();
fn delivery(
&self,
delivery_result: &DeliveryResult<'_>,
_delivery_opaque: Self::DeliveryOpaque,
) {
match delivery_result {
Ok(_message) => self.state.delivered(),
Err((error, message)) => self.state.failed(error, message),
}
}
}
#[cfg(feature = "rdkafka")]
struct DeliveryState {
sent: AtomicU64,
delivered: AtomicU64,
failed: AtomicU64,
first_error: Mutex<Option<String>>,
metrics: KafkaMetrics,
}
#[cfg(feature = "rdkafka")]
impl DeliveryState {
fn new(metrics: KafkaMetrics) -> Self {
Self {
sent: AtomicU64::new(0),
delivered: AtomicU64::new(0),
failed: AtomicU64::new(0),
first_error: Mutex::new(None),
metrics,
}
}
fn sent(&self) -> u64 {
let sent = self.sent.fetch_add(1, Ordering::Relaxed) + 1;
sent.saturating_sub(self.completed())
}
fn delivered(&self) {
self.delivered.fetch_add(1, Ordering::Relaxed);
self.metrics.produced();
self.metrics.producer_in_flight(self.in_flight());
}
fn failed(&self, error: &KafkaError, message: &rdkafka::message::BorrowedMessage<'_>) {
self.failed.fetch_add(1, Ordering::Relaxed);
self.metrics.delivery_failed();
self.metrics.producer_in_flight(self.in_flight());
if let Ok(mut slot) = self.first_error.lock()
&& slot.is_none()
{
*slot = Some(format!(
"{error}; topic={}, partition={}, offset={}",
message.topic(),
message.partition(),
message.offset()
));
}
}
fn completed(&self) -> u64 {
self.delivered.load(Ordering::Relaxed) + self.failed.load(Ordering::Relaxed)
}
fn in_flight(&self) -> u64 {
self.sent
.load(Ordering::Relaxed)
.saturating_sub(self.completed())
}
fn take_error(&self) -> Option<String> {
self.first_error
.lock()
.ok()
.and_then(|mut slot| slot.take())
}
}
#[cfg(feature = "rdkafka")]
fn to_base_record(record: &ProducerRecord) -> BaseRecord<'_, [u8], [u8], ()> {
let mut base_record = BaseRecord::<[u8], [u8], ()>::to(&record.topic);
if let Some(payload) = &record.payload {
base_record = base_record.payload(payload.as_ref());
}
if let Some(key) = &record.key {
base_record = base_record.key(key.as_ref());
}
if let Some(partition) = record.partition {
base_record = base_record.partition(partition);
}
if let Some(timestamp) = record.timestamp {
base_record = base_record.timestamp(timestamp);
}
if !record.headers.is_empty() {
base_record = base_record.headers(to_owned_headers(record));
}
base_record
}
#[cfg(feature = "rdkafka")]
fn is_queue_full(error: &KafkaError) -> bool {
matches!(
error,
KafkaError::MessageProduction(RDKafkaErrorCode::QueueFull)
)
}
#[cfg(feature = "rdkafka")]
fn to_owned_headers(record: &ProducerRecord) -> OwnedHeaders {
let mut headers = OwnedHeaders::new_with_capacity(record.headers.len());
for header in &record.headers {
headers = match &header.value {
Some(value) => headers.insert(Header {
key: &header.key,
value: Some(value.as_ref()),
}),
None => headers.insert(Header {
key: &header.key,
value: None::<&[u8]>,
}),
};
}
headers
}
#[cfg(feature = "rdkafka")]
#[allow(dead_code)]
fn _assert_client_config_send_sync(_: ClientConfig) {}