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