use std::sync::Arc;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::time::Duration;
use tokio::time::Instant;
use super::bounded::{
BoundedWorker, DrainHandle, DrainOutcome, DrainWitness, Recv, WorkerReceiver,
};
use crate::config::{AsyncOnOverflow, TraceStorageConfig, TraceStorageMode};
use crate::metrics;
use crate::storage::repositories::traces::{TraceCompletedRow, TraceResultRow, TraceSink};
#[derive(Debug)]
pub enum TracePersistenceTask {
StoreCompleted(TraceCompletedRow),
SetResult(TraceResultRow),
UpdateStatus {
id: String,
status: String,
error_message: Option<String>,
},
}
#[derive(Clone)]
pub struct TracePersistenceQueue {
queue: BoundedWorker<TracePersistenceTask>,
overflow_policy: AsyncOnOverflow,
overflow_block_timeout: Duration,
dropped_since_warn: Arc<AtomicUsize>,
last_warn_ms: Arc<AtomicU64>,
started: Instant,
}
const NEVER_WARNED: u64 = u64::MAX;
const OVERFLOW_WARN_INTERVAL_MS: u64 = 5_000;
impl TracePersistenceQueue {
pub fn disabled() -> Self {
Self {
queue: BoundedWorker::disabled(metrics::set_trace_persistence_queue_depth),
overflow_policy: AsyncOnOverflow::Drop,
overflow_block_timeout: Duration::ZERO,
dropped_since_warn: Arc::new(AtomicUsize::new(0)),
last_warn_ms: Arc::new(AtomicU64::new(NEVER_WARNED)),
started: Instant::now(),
}
}
pub async fn submit(&self, task: TracePersistenceTask) -> bool {
if self.queue.is_disabled() {
return false;
}
let accepted = match self.overflow_policy {
AsyncOnOverflow::Drop => self.queue.try_submit(task).is_ok(),
AsyncOnOverflow::Block => self
.queue
.submit_blocking(task, self.overflow_block_timeout)
.await
.is_ok(),
};
if !accepted {
metrics::record_trace_dropped("overflow");
self.dropped_since_warn.fetch_add(1, Ordering::Relaxed);
self.warn_if_window_elapsed();
}
accepted
}
fn warn_if_window_elapsed(&self) {
let elapsed = self.started.elapsed().as_millis() as u64;
let last = self.last_warn_ms.load(Ordering::Relaxed);
let due = last == NEVER_WARNED || elapsed.saturating_sub(last) >= OVERFLOW_WARN_INTERVAL_MS;
if !due
|| self
.last_warn_ms
.compare_exchange(last, elapsed, Ordering::Relaxed, Ordering::Relaxed)
.is_err()
{
return;
}
let dropped = self.dropped_since_warn.swap(0, Ordering::Relaxed);
tracing::warn!(
dropped,
window_ms = OVERFLOW_WARN_INTERVAL_MS,
"trace_persistence: queue full, dropping traces — the persistence workers cannot \
keep up with the request rate. Raise trace_storage.max_pending / batch_size, set \
trace_storage.async_on_overflow = \"block\" to slow producers instead, or use \
mode = \"sync\" so the request path cannot outrun the trace table"
);
}
}
pub struct PersistenceWorkerHandle {
drain: Option<DrainHandle<TracePersistenceTask>>,
joins: Vec<tokio::task::JoinHandle<()>>,
shutdown_timeout: Duration,
}
impl PersistenceWorkerHandle {
pub fn noop() -> Self {
Self {
drain: None,
joins: Vec::new(),
shutdown_timeout: Duration::ZERO,
}
}
pub async fn shutdown(self) {
let Some(drain) = self.drain else {
return;
};
match drain
.drain(self.joins, DrainWitness::TasksExit, self.shutdown_timeout)
.await
{
DrainOutcome::Drained => {}
DrainOutcome::WorkerPanicked => {
tracing::error!("Trace persistence worker panicked")
}
DrainOutcome::TimedOut { .. } => {
tracing::warn!("Trace persistence worker did not finish within shutdown timeout")
}
}
}
}
pub fn start(
tasks: &crate::runtime::TaskRegistry,
config: &TraceStorageConfig,
trace_repo: Arc<dyn TraceSink>,
) -> (TracePersistenceQueue, PersistenceWorkerHandle) {
let (worker_count, is_batch) = match config.mode {
TraceStorageMode::Async => (config.async_workers.max(1), false),
TraceStorageMode::Batch => (config.batch_workers.max(1), true),
TraceStorageMode::Sync | TraceStorageMode::Off => {
return (
TracePersistenceQueue::disabled(),
PersistenceWorkerHandle::noop(),
);
}
};
let per_worker_capacity = (config.max_pending.max(1) / worker_count).max(1);
let (queue, receivers) = BoundedWorker::<TracePersistenceTask>::new(
worker_count,
per_worker_capacity,
metrics::set_trace_persistence_queue_depth,
);
let drain = queue.drain_handle();
let guard = tasks.guard("trace_persistence", crate::runtime::Criticality::Required);
let mut joins = Vec::with_capacity(worker_count);
for rx in receivers {
let trace_repo = trace_repo.clone();
let batch_size = config.batch_size.max(1);
let flush_interval = Duration::from_millis(config.batch_flush_interval_ms.max(1));
joins.push(tokio::spawn(guard.clone().run(async move {
if is_batch {
run_batch_worker(rx, trace_repo, batch_size, flush_interval).await;
} else {
run_async_worker(rx, trace_repo).await;
}
})));
}
let handle = PersistenceWorkerHandle {
drain: Some(drain),
joins,
shutdown_timeout: Duration::from_secs(30),
};
(
TracePersistenceQueue {
queue,
overflow_policy: config.async_on_overflow,
overflow_block_timeout: Duration::from_millis(config.overflow_block_timeout_ms),
dropped_since_warn: Arc::new(AtomicUsize::new(0)),
last_warn_ms: Arc::new(AtomicU64::new(NEVER_WARNED)),
started: Instant::now(),
},
handle,
)
}
async fn run_async_worker(
mut rx: WorkerReceiver<TracePersistenceTask>,
trace_repo: Arc<dyn TraceSink>,
) {
while let Some(task) = rx.recv().await {
dispatch_one(&trace_repo, task).await;
}
}
async fn run_batch_worker(
mut rx: WorkerReceiver<TracePersistenceTask>,
trace_repo: Arc<dyn TraceSink>,
batch_size: usize,
flush_interval: Duration,
) {
let mut completed: Vec<TraceCompletedRow> = Vec::new();
let mut results: Vec<TraceResultRow> = Vec::new();
let mut deadline = Instant::now() + flush_interval;
loop {
let now = Instant::now();
let until = deadline.saturating_duration_since(now);
match rx.recv_timeout(until).await {
Recv::Item(task) => {
match task {
TracePersistenceTask::StoreCompleted(row) => completed.push(row),
TracePersistenceTask::SetResult(row) => results.push(row),
TracePersistenceTask::UpdateStatus {
id,
status,
error_message,
} => {
if let Err(e) = trace_repo
.update_status(&id, &status, error_message.as_deref())
.await
{
tracing::warn!(error = %e, "trace_persistence: update_status failed");
}
}
}
if completed.len() >= batch_size || results.len() >= batch_size {
flush_batches(&trace_repo, &mut completed, &mut results).await;
deadline = Instant::now() + flush_interval;
}
}
Recv::Closed => {
flush_batches(&trace_repo, &mut completed, &mut results).await;
return;
}
Recv::Elapsed => {
flush_batches(&trace_repo, &mut completed, &mut results).await;
deadline = Instant::now() + flush_interval;
}
}
}
}
const WRITE_RETRY_DELAYS: [Duration; 2] = [Duration::from_millis(50), Duration::from_millis(250)];
pub(super) async fn with_write_retries<F, Fut>(mut op: F) -> Result<(), crate::errors::OrionError>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Result<(), crate::errors::OrionError>>,
{
let mut last_err = None;
for (attempt, delay) in std::iter::once(Duration::ZERO)
.chain(WRITE_RETRY_DELAYS)
.enumerate()
{
if !delay.is_zero() {
tokio::time::sleep(delay).await;
tracing::debug!(attempt, "trace_persistence: retrying failed write");
}
match op().await {
Ok(()) => return Ok(()),
Err(e) => last_err = Some(e),
}
}
Err(last_err.expect("at least one attempt always runs"))
}
async fn dispatch_one(trace_repo: &Arc<dyn TraceSink>, task: TracePersistenceTask) {
let result = with_write_retries(|| async {
match &task {
TracePersistenceTask::StoreCompleted(row) => {
trace_repo.store_completed(row.as_view()).await.map(|_| ())
}
TracePersistenceTask::SetResult(row) => {
trace_repo
.set_result(
&row.id,
&row.result_json,
row.duration_ms,
row.task_trace_json.as_deref(),
)
.await
}
TracePersistenceTask::UpdateStatus {
id,
status,
error_message,
} => trace_repo
.update_status(id, status, error_message.as_deref())
.await
.map(|_| ()),
}
})
.await;
if let Err(e) = result {
crate::metrics::record_trace_persistence_failure();
tracing::warn!(error = %e, "trace_persistence: write failed after retries, dropping");
}
}
async fn flush_batches(
trace_repo: &Arc<dyn TraceSink>,
completed: &mut Vec<TraceCompletedRow>,
results: &mut Vec<TraceResultRow>,
) {
if !completed.is_empty() {
if let Err(e) = with_write_retries(|| async {
trace_repo
.store_completed_batch(completed)
.await
.map(|_| ())
})
.await
{
crate::metrics::record_trace_persistence_failure();
tracing::warn!(
error = %e,
dropped = completed.len(),
"trace_persistence: store_completed_batch failed after retries, dropping"
);
}
completed.clear();
}
if !results.is_empty() {
if let Err(e) =
with_write_retries(|| async { trace_repo.set_result_batch(results).await }).await
{
crate::metrics::record_trace_persistence_failure();
tracing::warn!(
error = %e,
dropped = results.len(),
"trace_persistence: set_result_batch failed after retries, dropping"
);
}
results.clear();
}
}