use super::breaker::BreakerSet;
use super::config::SinkPoolConfig;
use super::retry::Backoff;
use super::{EncodedChunk, SealedBatch, ShardWriter};
use crate::backpressure::InflightBudget;
use crate::checkpoint::AckSet;
use crate::error::{ErrorClass, SinkError};
use crate::metrics::{AttemptOutcome, FlushReason, SinkShardMetrics};
use std::collections::{HashMap, VecDeque};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::{Semaphore, mpsc, watch};
use tokio::task::{JoinError, JoinSet};
use tokio::time::Instant;
const ABORT_GRACE: Duration = Duration::from_millis(500);
const QUARANTINE_BACKSTOP_MIN: Duration = Duration::from_millis(100);
const QUARANTINE_BACKSTOP_MAX: Duration = Duration::from_secs(30);
fn quarantine_backstop(open_for: Duration) -> Duration {
open_for.clamp(QUARANTINE_BACKSTOP_MIN, QUARANTINE_BACKSTOP_MAX)
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub(crate) struct WorkerReport {
pub(crate) flushed: u64,
pub(crate) abandoned: u64,
}
impl WorkerReport {
pub(crate) fn absorb(&mut self, other: WorkerReport) {
self.flushed += other.flushed;
self.abandoned += other.abandoned;
}
}
struct Pending {
acks: AckSet,
rows: u64,
bytes: u64,
reason: FlushReason,
started: Instant,
oldest_ingest: std::time::Instant,
oldest_event_ms: i64,
}
struct WriteDone {
seq: u64,
written: bool,
}
struct Accumulator {
frames: Vec<bytes::Bytes>,
rows: u64,
bytes: u64,
acks: AckSet,
first_at: Option<Instant>,
oldest_ingest: Option<std::time::Instant>,
oldest_event_ms: i64,
}
impl Accumulator {
fn new() -> Self {
Accumulator {
frames: Vec::new(),
rows: 0,
bytes: 0,
acks: AckSet::new(),
first_at: None,
oldest_ingest: None,
oldest_event_ms: i64::MAX,
}
}
fn push(&mut self, chunk: EncodedChunk, now: Instant) {
self.first_at.get_or_insert(now);
self.rows += u64::from(chunk.rows);
self.bytes += chunk.frame.len() as u64;
self.frames.push(chunk.frame);
self.oldest_ingest = Some(match self.oldest_ingest {
Some(cur) => cur.min(chunk.oldest_ingest),
None => chunk.oldest_ingest,
});
self.oldest_event_ms = self.oldest_event_ms.min(chunk.oldest_event_ms);
self.acks.absorb(chunk.acks);
}
fn is_empty(&self) -> bool {
self.frames.is_empty()
}
}
pub(crate) struct ShardWorker<W: ShardWriter> {
pub(crate) shard: u32,
pub(crate) writer: Arc<W>,
pub(crate) endpoints: Arc<Vec<W::Endpoint>>,
pub(crate) rx: mpsc::Receiver<EncodedChunk>,
pub(crate) cfg: SinkPoolConfig,
pub(crate) budget: Arc<InflightBudget>,
pub(crate) metrics: Arc<SinkShardMetrics>,
pub(crate) drain_deadline: watch::Receiver<Option<Instant>>,
pub(crate) token_prefix: String,
}
struct Ledger {
pending: HashMap<u64, Pending>,
ids: HashMap<tokio::task::Id, u64>,
report: WorkerReport,
budget: Arc<InflightBudget>,
}
impl Drop for Ledger {
fn drop(&mut self) {
for p in self.pending.values() {
self.budget
.sub(usize::try_from(p.bytes).unwrap_or(usize::MAX));
}
}
}
impl<W: ShardWriter> ShardWorker<W> {
pub(crate) async fn run(mut self) -> WorkerReport {
let mut acc = Accumulator::new();
let mut ledger = Ledger {
pending: HashMap::new(),
ids: HashMap::new(),
report: WorkerReport::default(),
budget: Arc::clone(&self.budget),
};
let mut tasks: JoinSet<WriteDone> = JoinSet::new();
let semaphore = Arc::new(Semaphore::new(self.cfg.inflight.max_per_shard));
let breakers = Arc::new(Mutex::new(BreakerSet::new(
self.endpoints.len(),
self.cfg.breaker,
Arc::clone(&self.metrics),
)));
let mut drain_deadline = self.drain_deadline.clone();
let mut deadline_watch_live = true;
let mut seq: u64 = 0;
let mut recv_buf: Vec<EncodedChunk> = Vec::with_capacity(64);
let mut waiting: VecDeque<(u64, SealedBatch, Instant)> = VecDeque::new();
loop {
let linger_at = acc.first_at.map(|t| t + self.cfg.batch.linger);
tokio::select! {
biased;
Some(joined) = tasks.join_next_with_id(), if !tasks.is_empty() => {
self.handle_join(joined, &mut ledger);
}
permit = Arc::clone(&semaphore).acquire_owned(), if !waiting.is_empty() => {
let permit = permit.expect("sink semaphore closed");
self.launch_waiting(permit, &mut waiting, &mut tasks, &breakers, &mut ledger);
}
n = self.rx.recv_many(&mut recv_buf, 64), if waiting.is_empty() => {
if n == 0 {
break;
}
let now = Instant::now();
for chunk in recv_buf.drain(..) {
acc.push(chunk, now);
if let Some(reason) = self.seal_reason(&acc) {
self.dispatch(&mut acc, reason, &mut seq, &mut ledger, &mut tasks, &semaphore, &breakers, &mut waiting);
}
}
}
changed = drain_deadline.changed(), if deadline_watch_live => {
match changed {
Ok(()) if drain_deadline.borrow().is_some() => break,
Ok(()) => {}
Err(_) => deadline_watch_live = false,
}
}
() = tokio::time::sleep_until(linger_at.unwrap_or_else(Instant::now)), if linger_at.is_some() && waiting.is_empty() => {
self.dispatch(&mut acc, FlushReason::Linger, &mut seq, &mut ledger, &mut tasks, &semaphore, &breakers, &mut waiting);
}
}
}
while let Ok(chunk) = self.rx.try_recv() {
acc.push(chunk, Instant::now());
if let Some(reason) = self.seal_reason(&acc) {
self.dispatch(
&mut acc,
reason,
&mut seq,
&mut ledger,
&mut tasks,
&semaphore,
&breakers,
&mut waiting,
);
}
}
self.dispatch(
&mut acc,
FlushReason::Drain,
&mut seq,
&mut ledger,
&mut tasks,
&semaphore,
&breakers,
&mut waiting,
);
loop {
while !waiting.is_empty() {
let Ok(permit) = Arc::clone(&semaphore).try_acquire_owned() else {
break;
};
self.launch_waiting(permit, &mut waiting, &mut tasks, &breakers, &mut ledger);
}
if tasks.is_empty() && waiting.is_empty() {
break;
}
let deadline = *drain_deadline.borrow();
if deadline.is_none() && !deadline_watch_live {
self.sweep(&mut tasks, &mut waiting, &mut ledger).await;
break;
}
tokio::select! {
biased;
Some(joined) = tasks.join_next_with_id(), if !tasks.is_empty() => {
self.handle_join(joined, &mut ledger);
}
changed = drain_deadline.changed(), if deadline_watch_live => {
if changed.is_err() {
deadline_watch_live = false;
}
}
() = tokio::time::sleep_until(deadline.unwrap_or_else(Instant::now)), if deadline.is_some() => {
self.sweep(&mut tasks, &mut waiting, &mut ledger).await;
break;
}
}
}
ledger.report
}
fn seal_reason(&self, acc: &Accumulator) -> Option<FlushReason> {
if acc.rows >= self.cfg.batch.max_rows {
Some(FlushReason::Rows)
} else if acc.bytes >= self.cfg.batch.max_bytes {
Some(FlushReason::Bytes)
} else {
None
}
}
fn launch_waiting(
&self,
permit: tokio::sync::OwnedSemaphorePermit,
waiting: &mut VecDeque<(u64, SealedBatch, Instant)>,
tasks: &mut JoinSet<WriteDone>,
breakers: &Arc<Mutex<BreakerSet>>,
ledger: &mut Ledger,
) {
let (this_seq, batch, queued_at) = waiting
.pop_front()
.expect("callers guard on a non-empty `waiting`");
self.metrics.permit_waited(queued_at.elapsed());
self.spawn_write(batch, this_seq, permit, tasks, breakers, &mut ledger.ids);
}
async fn sweep(
&self,
tasks: &mut JoinSet<WriteDone>,
waiting: &mut VecDeque<(u64, SealedBatch, Instant)>,
ledger: &mut Ledger,
) {
if tokio::time::timeout(ABORT_GRACE, tasks.shutdown())
.await
.is_err()
{
tracing::error!(
shard = self.shard,
grace = ?ABORT_GRACE,
"sink write tasks did not abort within the grace period; abandoning without them"
);
}
waiting.clear();
let stranded: Vec<u64> = ledger.pending.keys().copied().collect();
for s in stranded {
self.abandon(s, ledger);
}
ledger.ids.clear();
}
fn handle_join(
&self,
joined: Result<(tokio::task::Id, WriteDone), JoinError>,
ledger: &mut Ledger,
) {
match joined {
Ok((id, WriteDone { seq, written })) => {
ledger.ids.remove(&id);
if written {
self.settle(seq, ledger);
} else {
self.abandon(seq, ledger);
}
}
Err(join_err) => {
let id = join_err.id();
tracing::error!(error = %join_err, "sink write task panicked");
match ledger.ids.remove(&id) {
Some(seq) => self.abandon(seq, ledger),
None => tracing::error!(
"panicked sink task had no ledger entry; batch already resolved"
),
}
}
}
}
fn settle(&self, seq: u64, ledger: &mut Ledger) {
let Some(p) = ledger.pending.remove(&seq) else {
return;
};
self.metrics
.flushed(p.reason, p.rows, p.bytes, p.started.elapsed());
self.metrics
.e2e_observed(p.oldest_ingest.elapsed(), p.oldest_event_ms);
self.budget
.sub(usize::try_from(p.bytes).unwrap_or(usize::MAX));
self.metrics.set_inflight(ledger.pending.len());
ledger.report.flushed += 1;
p.acks.deliver();
}
fn abandon(&self, seq: u64, ledger: &mut Ledger) {
let Some(p) = ledger.pending.remove(&seq) else {
return;
};
tracing::error!(
rows = p.rows,
bytes = p.bytes,
"abandoning sink batch; data will replay after restart"
);
drop(p.acks); self.metrics.abandoned(1);
self.budget
.sub(usize::try_from(p.bytes).unwrap_or(usize::MAX));
self.metrics.set_inflight(ledger.pending.len());
ledger.report.abandoned += 1;
}
fn seal(
&self,
acc: &mut Accumulator,
reason: FlushReason,
seq: &mut u64,
ledger: &mut Ledger,
) -> (u64, SealedBatch) {
let this_seq = *seq;
*seq += 1;
let full = std::mem::replace(acc, Accumulator::new());
let batch = SealedBatch {
frames: full.frames,
rows: full.rows,
bytes: full.bytes,
dedup_token: format!("{}{}", self.token_prefix, this_seq),
};
ledger.pending.insert(
this_seq,
Pending {
acks: full.acks,
rows: full.rows,
bytes: full.bytes,
reason,
started: Instant::now(),
oldest_ingest: full.oldest_ingest.unwrap_or_else(std::time::Instant::now),
oldest_event_ms: full.oldest_event_ms,
},
);
self.metrics.set_inflight(ledger.pending.len());
(this_seq, batch)
}
#[allow(clippy::too_many_arguments)]
fn spawn_write(
&self,
batch: SealedBatch,
this_seq: u64,
permit: tokio::sync::OwnedSemaphorePermit,
tasks: &mut JoinSet<WriteDone>,
breakers: &Arc<Mutex<BreakerSet>>,
ids: &mut HashMap<tokio::task::Id, u64>,
) {
let writer = Arc::clone(&self.writer);
let endpoints = Arc::clone(&self.endpoints);
let breakers = Arc::clone(breakers);
let metrics = Arc::clone(&self.metrics);
let retry = self.cfg.retry;
let backstop = quarantine_backstop(self.cfg.breaker.open_for);
let shard = self.shard;
let handle = tasks.spawn(async move {
let _permit = permit;
let mut backoff = Backoff::new(retry, this_seq);
let mut attempts: u32 = 0;
loop {
let pick = loop {
let now = Instant::now();
let (picked, probe_at) = {
let mut b = breakers.lock().expect("breaker lock");
let probe_at = b.next_probe_at(now);
let picked = match b.next_replica(now) {
Some(p) => Picked::Replica(p),
None => Picked::Park(b.subscribe()),
};
(picked, probe_at)
};
match picked {
Picked::Replica(p) => break p,
Picked::Park(mut wake) => {
let heartbeat = now + backstop;
let until = probe_at.map_or(heartbeat, |t| t.min(heartbeat));
tokio::select! {
biased;
changed = wake.changed() => {
if changed.is_err() {
tokio::time::sleep_until(until).await;
}
}
() = tokio::time::sleep_until(until) => {}
}
}
}
};
let replica = pick.replica;
let mut probe = pick
.probe
.map(|episode| ProbeGuard::new(&breakers, replica, episode));
attempts += 1;
let attempt_at = Instant::now();
let outcome = writer.write_batch(&endpoints[replica], &batch).await;
metrics.write_attempt(
if outcome.is_ok() {
AttemptOutcome::Ok
} else {
AttemptOutcome::Error
},
attempt_at.elapsed(),
);
if let Some(g) = probe.as_mut() {
g.disarm();
}
match outcome {
Ok(()) => {
let transition = breakers.lock().expect("breaker lock").on_success(replica);
if let Some(t) = transition {
t.log(shard);
}
return WriteDone {
seq: this_seq,
written: true,
};
}
Err(err) => {
let transition = breakers
.lock()
.expect("breaker lock")
.on_failure(replica, Instant::now());
if let Some(t) = transition {
t.log(shard);
}
let class = class_of(&err);
metrics.errors(class, 1);
metrics.replica_error(replica);
tracing::warn!(replica, attempts, error = %err, "sink write failed");
if class != ErrorClass::Retryable
|| (retry.max_attempts > 0 && attempts >= retry.max_attempts)
{
return WriteDone {
seq: this_seq,
written: false,
};
}
metrics.retries(1);
let delay = backoff.next_delay();
{
let _backoff = metrics.backing_off(this_seq, delay);
tokio::time::sleep(delay).await;
}
}
}
}
});
ids.insert(handle.id(), this_seq);
}
#[allow(clippy::too_many_arguments)]
fn dispatch(
&self,
acc: &mut Accumulator,
reason: FlushReason,
seq: &mut u64,
ledger: &mut Ledger,
tasks: &mut JoinSet<WriteDone>,
semaphore: &Arc<Semaphore>,
breakers: &Arc<Mutex<BreakerSet>>,
waiting: &mut VecDeque<(u64, SealedBatch, Instant)>,
) {
if acc.is_empty() {
return;
}
let (this_seq, batch) = self.seal(acc, reason, seq, ledger);
let queued_at = Instant::now();
if waiting.is_empty()
&& let Ok(permit) = Arc::clone(semaphore).try_acquire_owned()
{
self.metrics.permit_waited(queued_at.elapsed());
self.spawn_write(batch, this_seq, permit, tasks, breakers, &mut ledger.ids);
} else {
waiting.push_back((this_seq, batch, queued_at));
}
}
}
enum Picked {
Replica(super::breaker::Pick),
Park(watch::Receiver<u64>),
}
struct ProbeGuard {
breakers: Arc<Mutex<BreakerSet>>,
replica: usize,
episode: u64,
armed: bool,
}
impl ProbeGuard {
fn new(breakers: &Arc<Mutex<BreakerSet>>, replica: usize, episode: u64) -> Self {
ProbeGuard {
breakers: Arc::clone(breakers),
replica,
episode,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for ProbeGuard {
fn drop(&mut self) {
if !self.armed {
return;
}
self.breakers
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.release_probe(self.replica, self.episode);
}
}
fn class_of(err: &SinkError) -> ErrorClass {
match err {
SinkError::Client { class, .. } => *class,
#[allow(unreachable_patterns)]
_ => ErrorClass::Fatal,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_quarantine_heartbeat_is_clamped_at_both_ends() {
assert_eq!(
quarantine_backstop(Duration::ZERO),
QUARANTINE_BACKSTOP_MIN,
"a zero heartbeat is a spin, not a wait"
);
assert_eq!(
quarantine_backstop(Duration::from_secs(5)),
Duration::from_secs(5),
"the default is inside the range, so the clamp is the identity"
);
assert_eq!(
quarantine_backstop(Duration::from_secs(3600)),
QUARANTINE_BACKSTOP_MAX
);
assert_eq!(quarantine_backstop(Duration::MAX), QUARANTINE_BACKSTOP_MAX);
}
}