use super::{DriverEvent, ThreadControl};
use crate::admin::HealthState;
use crate::backpressure::{InflightBudget, Transition, WatermarkController};
use crate::checkpoint::AckRef;
use crate::error::{ErrorClass, FatalError, SourceError};
use crate::metrics::{BackpressureMetrics, SourceMetrics};
use crate::ops::{BlockReason, PushOutcome, RunnableChain};
use crate::record::RawPayload;
use crate::sink::ShardQueues;
use crate::source::{LaneId, PayloadBatch, SourceLane};
use crate::telemetry::RateLimit;
use std::panic::AssertUnwindSafe;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::{Duration, Instant};
static POLL_ERROR_WARN: RateLimit = RateLimit::new(5, Duration::from_secs(10));
#[derive(Clone, Debug)]
pub(crate) struct DriverParams {
pub thread: usize,
pub max_records: usize,
pub poll_timeout: Duration,
pub idle_flush: Duration,
pub blocked_retry: Duration,
pub queue_low_ratio: f64,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum DriverExit {
Completed,
Failed,
}
pub(crate) struct DriverContext<L> {
pub params: DriverParams,
pub control: crossbeam_channel::Receiver<ThreadControl<L>>,
pub events: crossbeam_channel::Sender<DriverEvent>,
pub chain: Box<dyn RunnableChain>,
pub bp: WatermarkController,
pub budget: Arc<InflightBudget>,
pub queues: Vec<ShardQueues>,
pub health: Arc<HealthState>,
pub bp_metrics: BackpressureMetrics,
pub source_metrics: SourceMetrics,
pub shutdown: Arc<AtomicBool>,
}
pub(crate) fn run_driver<L: SourceLane>(ctx: DriverContext<L>) -> DriverExit {
let DriverContext {
params,
control,
events,
mut chain,
mut bp,
budget,
queues,
health,
bp_metrics,
source_metrics,
shutdown,
} = ctx;
let mut lanes: Vec<L> = Vec::new();
let mut next_lane = 0usize;
let mut last_data = Instant::now();
let mut flushed_since_data = false;
let mut pause_started: Option<Instant> = None;
let mut empty_polls: usize = 0;
let mut parked: Option<ThreadControl<L>> = None;
loop {
health.heartbeat(params.thread);
while let Some(msg) = parked.take().or_else(|| control.try_recv().ok()) {
match msg {
ThreadControl::AddLane(lane) => lanes.push(lane),
ThreadControl::StopLanes {
lanes: stop,
barrier,
deadline,
} => {
let mut stopped = 0usize;
lanes.retain(|l| {
let goes = stop.contains(&l.id());
stopped += usize::from(goes);
!goes
});
if stopped > 0 {
flush_until(
chain.as_mut(),
deadline,
&mut bp,
&events,
&health,
params.thread,
);
}
for _ in 0..stopped {
barrier.arrive();
}
}
ThreadControl::FlushNow => {
match chain.flush() {
PushOutcome::Done => flushed_since_data = true,
PushOutcome::Blocked { .. } => bp.on_send_rejected(),
PushOutcome::Fatal(error) => {
let _ = events.send(DriverEvent::Fatal {
thread: params.thread,
error,
});
}
}
}
ThreadControl::DropLanes { lanes: drop } => {
lanes.retain(|l| !drop.contains(&l.id()));
}
ThreadControl::Shutdown { barrier, deadline } => {
flush_until(
chain.as_mut(),
deadline,
&mut bp,
&events,
&health,
params.thread,
);
lanes.clear();
drop(chain);
barrier.arrive();
return DriverExit::Completed;
}
}
}
let queues_low = queues.iter().all(|q| q.all_below(params.queue_low_ratio));
if let Some(t) = bp.tick(&budget, queues_low) {
let owned: Vec<LaneId> = lanes.iter().map(SourceLane::id).collect();
apply_transition(t, &owned, &events, &bp_metrics, &mut pause_started);
}
if bp.is_paused() {
std::thread::sleep(params.poll_timeout);
continue;
}
if lanes.is_empty() {
match control.recv_timeout(params.poll_timeout) {
Ok(msg) => parked = Some(msg),
Err(crossbeam_channel::RecvTimeoutError::Timeout) => {}
Err(crossbeam_channel::RecvTimeoutError::Disconnected) => {
std::thread::sleep(params.poll_timeout);
}
}
idle_flush(
chain.as_mut(),
&mut last_data,
&mut flushed_since_data,
params.idle_flush,
&mut bp,
&events,
params.thread,
);
continue;
}
next_lane %= lanes.len();
let lane_idx = next_lane;
next_lane += 1;
let lane_timeout = if empty_polls >= lanes.len() {
params.poll_timeout
} else {
Duration::ZERO
};
let owned_ids: Vec<LaneId> = lanes.iter().map(SourceLane::id).collect();
let fatal_reported = {
let poll_started = Instant::now();
let polled = lanes[lane_idx].poll(params.max_records, lane_timeout);
source_metrics.poll_duration(poll_started.elapsed());
let mut fatal_reported = false;
match polled {
Ok(Some(mut batch)) => {
empty_polls = 0;
last_data = Instant::now();
flushed_since_data = false;
let mut counting = CountingBatch::new(&mut batch);
let outcome = drive_batch(
chain.as_mut(),
&mut counting,
&mut bp,
&budget,
&queues,
¶ms,
&events,
&owned_ids,
&health,
&bp_metrics,
&mut pause_started,
&shutdown,
);
source_metrics.batch(counting.records, counting.bytes);
if let Err(error) = outcome {
let _ = events.send(DriverEvent::Fatal {
thread: params.thread,
error,
});
fatal_reported = true;
}
}
Ok(None) => {
empty_polls = empty_polls.saturating_add(1);
idle_flush(
chain.as_mut(),
&mut last_data,
&mut flushed_since_data,
params.idle_flush,
&mut bp,
&events,
params.thread,
);
}
Err(e) if is_fatal(&e) => {
let _ = events.send(DriverEvent::Fatal {
thread: params.thread,
error: FatalError {
component: format!("driver-{}", params.thread),
reason: format!("source poll failed: {e}"),
},
});
fatal_reported = true;
}
Err(e) => {
empty_polls = empty_polls.saturating_add(1);
crate::rate_limited_warn!(
POLL_ERROR_WARN,
thread = params.thread,
error = %e,
"retryable source poll error"
);
}
}
fatal_reported
};
if fatal_reported {
drop(chain);
return park_until_shutdown(&control, lanes, &health, params.thread);
}
}
}
fn is_fatal(e: &SourceError) -> bool {
let SourceError::Client { class, .. } = e;
*class == ErrorClass::Fatal
}
fn park_until_shutdown<L: SourceLane>(
control: &crossbeam_channel::Receiver<ThreadControl<L>>,
mut lanes: Vec<L>,
health: &HealthState,
thread: usize,
) -> DriverExit {
loop {
health.heartbeat(thread);
match control.recv_timeout(Duration::from_millis(50)) {
Ok(ThreadControl::AddLane(lane)) => drop(lane),
Ok(ThreadControl::StopLanes {
lanes: stop,
barrier,
..
}) => {
let mut stopped = 0usize;
lanes.retain(|l| {
let goes = stop.contains(&l.id());
stopped += usize::from(goes);
!goes
});
for _ in 0..stopped {
barrier.arrive();
}
}
Ok(ThreadControl::DropLanes { lanes: drop }) => {
lanes.retain(|l| !drop.contains(&l.id()));
}
Ok(ThreadControl::FlushNow) => {}
Ok(ThreadControl::Shutdown { barrier, .. }) => {
lanes.clear();
barrier.arrive();
return DriverExit::Failed;
}
Err(crossbeam_channel::RecvTimeoutError::Timeout) => {}
Err(crossbeam_channel::RecvTimeoutError::Disconnected) => return DriverExit::Failed,
}
}
}
#[expect(
clippy::too_many_arguments,
reason = "free function over disjoint driver-state borrows"
)]
fn drive_batch(
chain: &mut dyn RunnableChain,
batch: &mut dyn PayloadBatch<'_>,
bp: &mut WatermarkController,
budget: &InflightBudget,
queues: &[ShardQueues],
params: &DriverParams,
events: &crossbeam_channel::Sender<DriverEvent>,
owned: &[LaneId],
health: &HealthState,
bp_metrics: &BackpressureMetrics,
pause_started: &mut Option<Instant>,
shutdown: &AtomicBool,
) -> Result<(), FatalError> {
let ack: AckRef = batch.ack().clone();
let mut from = 0usize;
loop {
health.heartbeat(params.thread);
let pushed = std::panic::catch_unwind(AssertUnwindSafe(|| chain.push_batch(batch, from)));
match pushed {
Ok(PushOutcome::Done) => return Ok(()),
Ok(PushOutcome::Blocked { resume_at, reason }) => {
debug_assert!(resume_at >= from, "resume cursor must not go backwards");
from = resume_at;
if shutdown.load(Ordering::Relaxed) {
tracing::warn!(
thread = params.thread,
"shutdown during a blocked batch; abandoning it for replay"
);
ack.fail();
chain.abandon_batch();
return Ok(());
}
if reason == BlockReason::Capacity {
bp.on_send_rejected();
let queues_low = queues.iter().all(|q| q.all_below(params.queue_low_ratio));
if let Some(t) = bp.tick(budget, queues_low) {
apply_transition(t, owned, events, bp_metrics, pause_started);
}
}
std::thread::sleep(params.blocked_retry);
}
Ok(PushOutcome::Fatal(error)) => {
ack.fail();
return Err(error);
}
Err(panic) => {
ack.fail();
return Err(FatalError {
component: format!("driver-{}", params.thread),
reason: format!("operator chain panicked: {}", panic_message(panic.as_ref())),
});
}
}
}
}
struct CountingBatch<'a, 'buf> {
inner: &'a mut dyn PayloadBatch<'buf>,
records: u64,
bytes: u64,
}
impl<'a, 'buf> CountingBatch<'a, 'buf> {
fn new(inner: &'a mut dyn PayloadBatch<'buf>) -> Self {
CountingBatch {
inner,
records: 0,
bytes: 0,
}
}
}
impl<'buf> PayloadBatch<'buf> for CountingBatch<'_, 'buf> {
fn next_payload(&mut self) -> Option<RawPayload<'buf>> {
let payload = self.inner.next_payload()?;
self.records += 1;
self.bytes +=
payload.bytes.len() as u64 + payload.key.map(<[u8]>::len).unwrap_or_default() as u64;
Some(payload)
}
fn ack(&self) -> &AckRef {
self.inner.ack()
}
}
fn flush_until(
chain: &mut dyn RunnableChain,
deadline: Instant,
bp: &mut WatermarkController,
events: &crossbeam_channel::Sender<DriverEvent>,
health: &HealthState,
thread: usize,
) {
loop {
health.heartbeat(thread);
match chain.flush() {
PushOutcome::Done => return,
PushOutcome::Blocked { .. } => {
bp.on_send_rejected();
if Instant::now() >= deadline {
tracing::error!(
thread,
"drain deadline exceeded with the chain still blocked; \
abandoning parked records for replay"
);
return;
}
std::thread::sleep(Duration::from_millis(2));
}
PushOutcome::Fatal(error) => {
let _ = events.send(DriverEvent::Fatal { thread, error });
return;
}
}
}
}
fn idle_flush(
chain: &mut dyn RunnableChain,
last_data: &mut Instant,
flushed_since_data: &mut bool,
after: Duration,
bp: &mut WatermarkController,
events: &crossbeam_channel::Sender<DriverEvent>,
thread: usize,
) {
if *flushed_since_data || last_data.elapsed() < after {
return;
}
match chain.flush() {
PushOutcome::Done => *flushed_since_data = true,
PushOutcome::Blocked { .. } => {
bp.on_send_rejected();
}
PushOutcome::Fatal(error) => {
let _ = events.send(DriverEvent::Fatal { thread, error });
}
}
}
fn apply_transition(
t: Transition,
owned: &[LaneId],
events: &crossbeam_channel::Sender<DriverEvent>,
bp_metrics: &BackpressureMetrics,
pause_started: &mut Option<Instant>,
) {
match t {
Transition::Pause => {
*pause_started = Some(Instant::now());
bp_metrics.pause_started();
let _ = events.send(DriverEvent::PauseLanes {
lanes: owned.to_vec(),
});
}
Transition::Resume => {
let paused_for = pause_started
.take()
.map(|s| s.elapsed())
.unwrap_or_default();
bp_metrics.pause_ended(paused_for);
let _ = events.send(DriverEvent::ResumeLanes {
lanes: owned.to_vec(),
});
}
}
}
fn panic_message(panic: &(dyn std::any::Any + Send)) -> String {
if let Some(s) = panic.downcast_ref::<&str>() {
(*s).to_string()
} else if let Some(s) = panic.downcast_ref::<String>() {
s.clone()
} else {
"non-string panic payload".to_string()
}
}