datum-mq 0.10.11

Kafka sources and sinks for Datum streams, with native and rdkafka backends
Documentation
use std::{
    sync::{
        Arc,
        atomic::{AtomicBool, Ordering},
        mpsc::{Receiver, RecvTimeoutError, SyncSender, TrySendError, sync_channel},
    },
    thread,
    time::{Duration, Instant},
};

use datum::{Source, StreamError};
use tokio::runtime::Builder;

use crate::native::{
    KafkaClientError, KafkaClientResult,
    client::{NativeCommitPolicy, NativeKafkaConsumer, NativeKafkaConsumerConfig},
    model::{KafkaPayloadBatch, LocalCommitState, NativeKafkaMetrics, NativeKafkaMetricsSnapshot},
    profile::{self, ProfileBucket},
};

type SourceMessage = KafkaClientResult<KafkaPayloadBatch>;

pub struct NativeKafkaSource;

#[derive(Debug, Clone)]
pub struct NativeKafkaControl {
    state: Arc<ControlState>,
}

#[derive(Debug)]
struct ControlState {
    shutdown: AtomicBool,
    draining: AtomicBool,
    commit_state: Arc<LocalCommitState>,
    metrics: NativeKafkaMetrics,
}

impl NativeKafkaControl {
    fn new(commit_state: Arc<LocalCommitState>, metrics: NativeKafkaMetrics) -> Self {
        Self {
            state: Arc::new(ControlState {
                shutdown: AtomicBool::new(false),
                draining: AtomicBool::new(false),
                commit_state,
                metrics,
            }),
        }
    }

    pub fn drain_and_shutdown(&self, timeout: Duration) -> KafkaClientResult<()> {
        self.state.draining.store(true, Ordering::SeqCst);
        if self.state.commit_state.wait_for_all_committed(timeout) {
            self.state.shutdown.store(true, Ordering::SeqCst);
            Ok(())
        } else {
            Err(KafkaClientError::protocol(
                "native Kafka source drain timed out",
            ))
        }
    }

    pub fn shutdown_now(&self) {
        self.state.shutdown.store(true, Ordering::SeqCst);
    }

    #[must_use]
    pub fn outstanding(&self) -> u64 {
        self.state.commit_state.outstanding()
    }

    #[must_use]
    pub fn committed_offsets_sum(&self) -> u64 {
        self.state.commit_state.committed_offsets_sum()
    }

    #[must_use]
    pub fn metrics(&self) -> NativeKafkaMetricsSnapshot {
        self.state.metrics.snapshot()
    }
}

impl NativeKafkaSource {
    #[must_use]
    pub fn payload_batches(
        mut config: NativeKafkaConsumerConfig,
    ) -> Source<KafkaPayloadBatch, NativeKafkaControl> {
        let commit_state = Arc::new(LocalCommitState::default());
        let metrics = NativeKafkaMetrics::default();
        config.metrics = metrics.clone();
        let control = NativeKafkaControl::new(Arc::clone(&commit_state), metrics);
        Source::unfold_resource(
            {
                let control = control.clone();
                move || {
                    NativeKafkaSourceResource::start(
                        config.clone(),
                        control.clone(),
                        Arc::clone(&commit_state),
                    )
                    .map_err(StreamError::from)
                }
            },
            NativeKafkaSourceResource::read_next,
            NativeKafkaSourceResource::close,
        )
        .map_materialized_value(move |_| control.clone())
    }
}

struct NativeKafkaSourceResource {
    control: NativeKafkaControl,
    receiver: Receiver<SourceMessage>,
    worker: Option<thread::JoinHandle<()>>,
}

impl NativeKafkaSourceResource {
    fn start(
        config: NativeKafkaConsumerConfig,
        control: NativeKafkaControl,
        commit_state: Arc<LocalCommitState>,
    ) -> KafkaClientResult<Self> {
        let (sender, receiver) = sync_channel(1);
        let worker_control = control.clone();
        let worker = thread::Builder::new()
            .name("datum-kafka-native-source".to_owned())
            .spawn(move || fetch_worker(config, worker_control, commit_state, sender))
            .map_err(|error| {
                KafkaClientError::protocol(format!(
                    "failed to spawn native Kafka source worker: {error}"
                ))
            })?;
        Ok(Self {
            control,
            receiver,
            worker: Some(worker),
        })
    }

    fn read_next(&mut self) -> datum::StreamResult<Option<KafkaPayloadBatch>> {
        loop {
            if self.control.state.shutdown.load(Ordering::SeqCst) {
                return Ok(None);
            }
            match self.receiver.recv_timeout(Duration::from_millis(10)) {
                Ok(Ok(batch)) => {
                    return Ok(Some(profile::measure(ProfileBucket::SourceEmit, || batch)));
                }
                Ok(Err(error)) => return Err(StreamError::Failed(error.to_string())),
                Err(RecvTimeoutError::Timeout) => continue,
                Err(RecvTimeoutError::Disconnected) => return Ok(None),
            }
        }
    }

    fn close(mut self) -> datum::StreamResult<()> {
        if self.control.outstanding() == 0 {
            let _ = self.control.drain_and_shutdown(Duration::from_secs(30));
        } else {
            self.control.shutdown_now();
        }
        drop(self.receiver);
        if let Some(worker) = self.worker.take() {
            worker.join().map_err(|_| {
                StreamError::Failed("native Kafka source worker panicked".to_owned())
            })?;
        }
        Ok(())
    }
}

fn fetch_worker(
    config: NativeKafkaConsumerConfig,
    control: NativeKafkaControl,
    commit_state: Arc<LocalCommitState>,
    sender: SyncSender<SourceMessage>,
) {
    let runtime = match Builder::new_current_thread().enable_all().build() {
        Ok(runtime) => runtime,
        Err(error) => {
            let _ = sender.send(Err(KafkaClientError::protocol(format!(
                "failed to start native Kafka Tokio runtime: {error}"
            ))));
            return;
        }
    };
    runtime.block_on(async move {
        let mut consumer = match NativeKafkaConsumer::connect(config).await {
            Ok(consumer) => consumer,
            Err(error) => {
                let _ = sender.send(Err(error));
                return;
            }
        };
        let mut last_commit_flush = Instant::now();
        loop {
            if control.state.shutdown.load(Ordering::SeqCst) {
                break;
            }
            if consumer.group_enabled()
                && let Err(error) = consumer.poll_group(&commit_state).await
            {
                let _ = sender.send(Err(error));
                break;
            }
            if let Err(error) = flush_commits_if_due(
                &mut consumer,
                &commit_state,
                &control.state.metrics,
                false,
                &mut last_commit_flush,
            )
            .await
            {
                let _ = sender.send(Err(error));
                break;
            }
            if control.state.draining.load(Ordering::SeqCst)
                && control.state.commit_state.outstanding() == 0
            {
                match flush_commits_if_due(
                    &mut consumer,
                    &commit_state,
                    &control.state.metrics,
                    true,
                    &mut last_commit_flush,
                )
                .await
                {
                    Ok(()) if commit_state.uncommitted() == 0 => break,
                    Ok(()) => {
                        tokio::time::sleep(Duration::from_millis(1)).await;
                        continue;
                    }
                    Err(error) => {
                        let _ = sender.send(Err(error));
                        break;
                    }
                }
            }
            if consumer.group_enabled() {
                consumer.refresh_pauses(&commit_state);
            }
            if control.state.commit_state.outstanding() as usize >= consumer.fetch.high_watermark {
                tokio::time::sleep(Duration::from_millis(1)).await;
                continue;
            }
            match fetch_or_stop(&mut consumer, &control).await {
                Ok(FetchStep::Batch(Some(batch))) => {
                    let high_watermark = batch
                        .watermarks()
                        .iter()
                        .map(|watermark| watermark.offset - 1)
                        .max()
                        .unwrap_or(0);
                    let records = batch.len() as u64;
                    let batch = batch.with_commit_state(Arc::clone(&commit_state));
                    control.state.metrics.emitted_batch(
                        records,
                        commit_state.outstanding(),
                        high_watermark,
                    );
                    if !send_batch(&mut consumer, &control, &commit_state, &sender, batch).await {
                        break;
                    }
                }
                Ok(FetchStep::Batch(None)) => tokio::time::sleep(Duration::from_millis(1)).await,
                Ok(FetchStep::Stop) => {
                    if let Err(error) = flush_commits_if_due(
                        &mut consumer,
                        &commit_state,
                        &control.state.metrics,
                        true,
                        &mut last_commit_flush,
                    )
                    .await
                    {
                        let _ = sender.send(Err(error));
                    }
                    break;
                }
                Err(error) => {
                    let _ = sender.send(Err(error));
                    break;
                }
            }
        }
        let _ = consumer.leave_group().await;
    });
}

async fn send_batch(
    consumer: &mut NativeKafkaConsumer,
    control: &NativeKafkaControl,
    commit_state: &LocalCommitState,
    sender: &SyncSender<SourceMessage>,
    batch: KafkaPayloadBatch,
) -> bool {
    let mut message = Ok(batch);
    let mut stop_after_send = false;
    loop {
        if control.state.shutdown.load(Ordering::SeqCst) {
            return false;
        }
        match sender.try_send(message) {
            Ok(()) => return !stop_after_send,
            Err(TrySendError::Full(next)) => {
                message = next;
                if consumer.group_enabled()
                    && let Err(error) = consumer.poll_group(commit_state).await
                {
                    message = Err(error);
                    stop_after_send = true;
                }
                tokio::time::sleep(Duration::from_millis(1)).await;
            }
            Err(TrySendError::Disconnected(_)) => return false,
        }
    }
}

async fn flush_commits_if_due(
    consumer: &mut NativeKafkaConsumer,
    commit_state: &LocalCommitState,
    metrics: &NativeKafkaMetrics,
    force: bool,
    last_commit_flush: &mut Instant,
) -> KafkaClientResult<()> {
    let due_by_count = commit_state.uncommitted() >= consumer.commit_batch_size() as u64;
    let due_by_time = last_commit_flush.elapsed() >= consumer.commit_interval();
    if !force && !due_by_count && !due_by_time {
        return Ok(());
    }
    let commits = commit_state.due_commits();
    if commits.is_empty() {
        *last_commit_flush = Instant::now();
        return Ok(());
    }
    profile::measure_async(ProfileBucket::CommitBookkeeping, async {
        if consumer.commit_policy() == NativeCommitPolicy::Manual {
            consumer.commit_offsets(&commits).await?;
        }
        commit_state.mark_committed(&commits);
        metrics.committed(
            commits.len() as u64,
            commit_state.outstanding(),
            commits
                .iter()
                .map(|commit| commit.offset - 1)
                .max()
                .unwrap_or(0),
        );
        *last_commit_flush = Instant::now();
        Ok(())
    })
    .await
}

enum FetchStep {
    Batch(Option<KafkaPayloadBatch>),
    Stop,
}

async fn fetch_or_stop(
    consumer: &mut NativeKafkaConsumer,
    control: &NativeKafkaControl,
) -> KafkaClientResult<FetchStep> {
    tokio::select! {
        result = consumer.fetch_batch() => result.map(FetchStep::Batch),
        () = wait_for_stop(control) => Ok(FetchStep::Stop),
    }
}

async fn wait_for_stop(control: &NativeKafkaControl) {
    loop {
        if control.state.shutdown.load(Ordering::SeqCst) {
            return;
        }
        if control.state.draining.load(Ordering::SeqCst)
            && control.state.commit_state.outstanding() == 0
        {
            return;
        }
        tokio::time::sleep(Duration::from_millis(1)).await;
    }
}

impl From<KafkaClientError> for StreamError {
    fn from(error: KafkaClientError) -> Self {
        StreamError::Failed(error.to_string())
    }
}