use std::sync::Arc;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::mpsc;
use tokio::time::Instant;
use crate::config::{AsyncOnOverflow, TraceStorageConfig, TraceStorageMode};
use crate::metrics;
use crate::storage::repositories::traces::{TraceCompletedRow, TraceRepository, TraceResultRow};
#[derive(Debug)]
pub enum TracePersistenceTask {
StoreCompleted(TraceCompletedRow),
SetResult(TraceResultRow),
UpdateStatus {
id: String,
status: String,
error_message: Option<String>,
},
}
#[derive(Clone)]
pub struct TracePersistenceQueue {
senders: Vec<mpsc::Sender<TracePersistenceTask>>,
next: Arc<AtomicUsize>,
pending: Arc<AtomicUsize>,
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 {
senders: Vec::new(),
next: Arc::new(AtomicUsize::new(0)),
pending: Arc::new(AtomicUsize::new(0)),
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.senders.is_empty() {
return false;
}
let start = self.next.fetch_add(1, Ordering::Relaxed);
let send_result = match self.overflow_policy {
AsyncOnOverflow::Drop => {
let mut task = task;
let mut outcome = Err("full");
for i in 0..self.senders.len() {
let sender = &self.senders[(start + i) % self.senders.len()];
match sender.try_send(task) {
Ok(()) => {
outcome = Ok(());
break;
}
Err(mpsc::error::TrySendError::Full(t))
| Err(mpsc::error::TrySendError::Closed(t)) => task = t,
}
}
outcome
}
AsyncOnOverflow::Block => {
let sender = &self.senders[start % self.senders.len()];
match tokio::time::timeout(self.overflow_block_timeout, sender.send(task)).await {
Ok(Ok(())) => Ok(()),
Ok(Err(_)) => Err("closed"),
Err(_) => Err("timeout"),
}
}
};
match send_result {
Ok(()) => {
let n = self.pending.fetch_add(1, Ordering::Relaxed) + 1;
metrics::set_trace_persistence_queue_depth(n as f64);
true
}
Err(_) => {
metrics::record_trace_dropped("overflow");
self.dropped_since_warn.fetch_add(1, Ordering::Relaxed);
self.warn_if_window_elapsed();
false
}
}
}
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 {
_senders: Vec<mpsc::Sender<TracePersistenceTask>>,
join: Vec<tokio::task::JoinHandle<()>>,
shutdown_timeout: Duration,
}
impl PersistenceWorkerHandle {
pub fn noop() -> Self {
Self {
_senders: Vec::new(),
join: Vec::new(),
shutdown_timeout: Duration::ZERO,
}
}
pub async fn shutdown(self) {
drop(self._senders);
if self.join.is_empty() {
return;
}
let deadline = tokio::time::Instant::now() + self.shutdown_timeout;
for handle in self.join {
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
handle.abort();
continue;
}
if tokio::time::timeout(remaining, handle).await.is_err() {
tracing::warn!("Trace persistence worker did not finish within shutdown timeout");
}
}
}
}
pub fn start(
config: &TraceStorageConfig,
trace_repo: Arc<dyn TraceRepository>,
) -> (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 pending = Arc::new(AtomicUsize::new(0));
let mut senders = Vec::with_capacity(worker_count);
let mut join = Vec::with_capacity(worker_count);
for _ in 0..worker_count {
let (tx, rx) = mpsc::channel::<TracePersistenceTask>(per_worker_capacity);
senders.push(tx);
let pending = pending.clone();
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));
join.push(tokio::spawn(async move {
if is_batch {
run_batch_worker(rx, pending, trace_repo, batch_size, flush_interval).await;
} else {
run_async_worker(rx, pending, trace_repo).await;
}
}));
}
let queue = TracePersistenceQueue {
senders: senders.clone(),
next: Arc::new(AtomicUsize::new(0)),
pending,
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(),
};
let handle = PersistenceWorkerHandle {
_senders: senders,
join,
shutdown_timeout: Duration::from_secs(30),
};
(queue, handle)
}
async fn run_async_worker(
mut rx: mpsc::Receiver<TracePersistenceTask>,
pending: Arc<AtomicUsize>,
trace_repo: Arc<dyn TraceRepository>,
) {
while let Some(task) = rx.recv().await {
let n = pending.fetch_sub(1, Ordering::Relaxed).saturating_sub(1);
metrics::set_trace_persistence_queue_depth(n as f64);
dispatch_one(&trace_repo, task).await;
}
}
async fn run_batch_worker(
mut rx: mpsc::Receiver<TracePersistenceTask>,
pending: Arc<AtomicUsize>,
trace_repo: Arc<dyn TraceRepository>,
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);
let recv = tokio::time::timeout(until, rx.recv()).await;
match recv {
Ok(Some(task)) => {
let n = pending.fetch_sub(1, Ordering::Relaxed).saturating_sub(1);
metrics::set_trace_persistence_queue_depth(n as f64);
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;
}
}
Ok(None) => {
flush_batches(&trace_repo, &mut completed, &mut results).await;
return;
}
Err(_) => {
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 TraceRepository>, task: TracePersistenceTask) {
let result = with_write_retries(|| async {
match &task {
TracePersistenceTask::StoreCompleted(row) => trace_repo
.store_completed(
&row.channel,
row.channel_id.as_deref(),
&row.mode,
row.input_json.as_deref(),
&row.result_json,
row.duration_ms,
row.task_trace_json.as_deref(),
)
.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 TraceRepository>,
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();
}
}