datum-mq 0.10.2

Kafka sources and sinks for Datum streams, backed by rdkafka
Documentation
use std::{
    sync::{
        Arc, Mutex,
        atomic::{AtomicU64, Ordering},
    },
    time::{Duration, Instant},
};

use datum::{Flow, Keep, NotUsed, Sink, StreamCompletion, StreamError};
use rdkafka::{
    ClientConfig, ClientContext,
    error::{KafkaError, RDKafkaErrorCode},
    message::{Header, Message, OwnedHeaders},
    producer::{BaseProducer, BaseRecord, DeliveryResult, Producer, ProducerContext},
    util::Timeout,
};

use crate::{KafkaMetrics, KafkaProducerSettings, MqError, MqResult, ProducerRecord};

const PRODUCER_BATCH_SIZE: usize = 256;
const PRODUCER_POLL_SLICE: Duration = Duration::from_millis(1);

/// Kafka producer sink entry points.
pub struct KafkaSink;

impl KafkaSink {
    /// Creates a producer sink.
    ///
    /// The input record carries its target topic. The sink completes each element
    /// by enqueuing bounded 256-record batches into a single `BaseProducer`
    /// owner. Delivery callbacks are polled by that owner and failures fail the
    /// stream on the next poll; call [`KafkaProducerControl::drain_and_shutdown`]
    /// at shutdown to wait for final acknowledgements before closing. Idempotence
    /// is enabled by [`KafkaProducerSettings::new`] unless callers override raw
    /// properties.
    #[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> {
        Sink::setup(move |_materializer, _attributes| {
            let control = KafkaProducerControl::new(settings.clone())
                .expect("Kafka producer configuration must be valid at materialization");
            let runner = Arc::clone(&control.state.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()
            })
        })
    }
}

/// Materialized control for a Kafka producer sink.
#[derive(Clone)]
pub struct KafkaProducerControl {
    state: Arc<ProducerControlState>,
}

struct ProducerControlState {
    runner: Arc<Mutex<ProducerRunner>>,
    metrics: KafkaMetrics,
    completion: Mutex<Option<StreamCompletion<NotUsed>>>,
}

impl KafkaProducerControl {
    fn new(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 {
                runner: Arc::new(Mutex::new(ProducerRunner {
                    producer,
                    delivery,
                    metrics: metrics.clone(),
                    settings,
                    enqueued_since_poll: 0,
                })),
                metrics,
                completion: Mutex::new(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()
    }

    /// Waits for accepted input records to receive delivery acknowledgements, then flushes.
    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();
        if let Some(completion) = completion {
            completion
                .wait()
                .map_err(|error| MqError::Failed(error.to_string()))?;
        }
        self.flush()
    }

    pub fn flush(&self) -> MqResult<()> {
        self.state
            .runner
            .lock()
            .map_err(|_| MqError::Failed("Kafka producer lock poisoned".to_owned()))?
            .finish()
    }
}

struct ProducerRunner {
    producer: BaseProducer<DeliveryContext>,
    delivery: Arc<DeliveryState>,
    metrics: KafkaMetrics,
    settings: KafkaProducerSettings,
    enqueued_since_poll: usize,
}

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)]
struct DeliveryContext {
    state: Arc<DeliveryState>,
}

impl ClientContext for DeliveryContext {}

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),
        }
    }
}

struct DeliveryState {
    sent: AtomicU64,
    delivered: AtomicU64,
    failed: AtomicU64,
    first_error: Mutex<Option<String>>,
    metrics: KafkaMetrics,
}

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())
    }
}

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
}

fn is_queue_full(error: &KafkaError) -> bool {
    matches!(
        error,
        KafkaError::MessageProduction(RDKafkaErrorCode::QueueFull)
    )
}

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
}

#[allow(dead_code)]
fn _assert_client_config_send_sync(_: ClientConfig) {}