use std::sync::Arc;
use std::sync::atomic::Ordering;
use async_trait::async_trait;
use crate::support::io::LogSink;
use crate::{InklogError, LogRecord, Metrics};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SinkWriteOutcome {
Written,
Rejected,
Failed,
}
pub trait SinkRateLimit: Send + Sync {
fn try_acquire(&self, record: &LogRecord) -> bool;
fn report(&self, _record: &LogRecord, _outcome: SinkWriteOutcome) {}
fn name(&self) -> &str {
"sink-rate-limit"
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct NoOpRateLimit;
impl SinkRateLimit for NoOpRateLimit {
fn try_acquire(&self, _record: &LogRecord) -> bool {
true
}
}
pub struct TokenBucketRateLimit {
tokens: std::sync::atomic::AtomicU64,
capacity: u64,
name: String,
}
impl TokenBucketRateLimit {
pub fn new(capacity: u64) -> Self {
Self {
tokens: std::sync::atomic::AtomicU64::new(capacity),
capacity,
name: format!("token-bucket({capacity})"),
}
}
pub fn available(&self) -> u64 {
self.tokens.load(Ordering::Relaxed)
}
pub fn capacity(&self) -> u64 {
self.capacity
}
}
impl SinkRateLimit for TokenBucketRateLimit {
fn try_acquire(&self, _record: &LogRecord) -> bool {
let mut current = self.tokens.load(Ordering::Relaxed);
loop {
if current == 0 {
return false;
}
match self.tokens.compare_exchange_weak(
current,
current - 1,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return true,
Err(actual) => current = actual,
}
}
}
fn report(&self, _record: &LogRecord, outcome: SinkWriteOutcome) {
if outcome == SinkWriteOutcome::Failed {
let mut current = self.tokens.load(Ordering::Relaxed);
while current < self.capacity {
match self.tokens.compare_exchange_weak(
current,
current + 1,
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => return,
Err(actual) => current = actual,
}
}
}
}
fn name(&self) -> &str {
&self.name
}
}
pub struct RateLimitedSink {
inner: Arc<dyn LogSink>,
limiter: Arc<dyn SinkRateLimit>,
metrics: Option<Arc<Metrics>>,
}
impl RateLimitedSink {
pub fn new(inner: Arc<dyn LogSink>, limiter: Arc<dyn SinkRateLimit>) -> Self {
Self {
inner,
limiter,
metrics: None,
}
}
pub fn with_metrics(mut self, metrics: Arc<Metrics>) -> Self {
self.metrics = Some(metrics);
self
}
}
#[async_trait]
impl LogSink for RateLimitedSink {
async fn write(&self, record: &LogRecord) -> Result<(), InklogError> {
if !self.limiter.try_acquire(record) {
if let Some(ref metrics) = self.metrics {
metrics.inc_logs_dropped();
}
self.limiter.report(record, SinkWriteOutcome::Rejected);
return Ok(());
}
match self.inner.write(record).await {
Ok(()) => {
self.limiter.report(record, SinkWriteOutcome::Written);
Ok(())
}
Err(e) => {
self.limiter.report(record, SinkWriteOutcome::Failed);
Err(e)
}
}
}
async fn flush(&self) -> Result<(), InklogError> {
self.inner.flush().await
}
fn is_healthy(&self) -> bool {
self.inner.is_healthy()
}
async fn shutdown(&self) -> Result<(), InklogError> {
self.inner.shutdown().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use parking_lot::Mutex;
fn record(level: tracing::Level, target: &str) -> LogRecord {
LogRecord::new(level, target.to_string(), "msg".to_string())
}
#[test]
fn test_noop_rate_limit_always_allows() {
let limiter = NoOpRateLimit;
for i in 0..1000 {
assert!(
limiter.try_acquire(&record(tracing::Level::INFO, &format!("t{i}"))),
"NoOp must always allow"
);
}
assert_eq!(limiter.name(), "sink-rate-limit");
}
#[test]
fn test_token_bucket_budget_and_refund() {
let limiter = TokenBucketRateLimit::new(3);
assert!(limiter.try_acquire(&record(tracing::Level::INFO, "t")));
assert!(limiter.try_acquire(&record(tracing::Level::INFO, "t")));
assert!(limiter.try_acquire(&record(tracing::Level::INFO, "t")));
assert_eq!(limiter.available(), 0);
assert!(
!limiter.try_acquire(&record(tracing::Level::ERROR, "t")),
"exhausted budget must reject (even severe records; per-target policy is upper-layer)"
);
limiter.report(&record(tracing::Level::INFO, "t"), SinkWriteOutcome::Failed);
assert_eq!(limiter.available(), 1);
assert!(limiter.try_acquire(&record(tracing::Level::INFO, "t")));
}
struct CollectingSink {
messages: Mutex<Vec<String>>,
fail_once: std::sync::atomic::AtomicBool,
}
#[async_trait]
impl LogSink for CollectingSink {
async fn write(&self, record: &LogRecord) -> Result<(), InklogError> {
if self.fail_once.swap(false, Ordering::Relaxed) {
return Err(InklogError::ConfigError("simulated failure".to_string()));
}
self.messages.lock().push(record.message.clone());
Ok(())
}
async fn flush(&self) -> Result<(), InklogError> {
Ok(())
}
async fn shutdown(&self) -> Result<(), InklogError> {
Ok(())
}
}
#[tokio::test]
async fn test_rate_limited_sink_blocks_writes_and_counts_dropped() {
let inner = Arc::new(CollectingSink {
messages: Mutex::new(Vec::new()),
fail_once: std::sync::atomic::AtomicBool::new(false),
});
let metrics = Arc::new(Metrics::new());
let sink = Arc::new(
RateLimitedSink::new(
inner.clone() as Arc<dyn LogSink>,
Arc::new(TokenBucketRateLimit::new(2)),
)
.with_metrics(metrics.clone()),
);
sink.write(&record(tracing::Level::INFO, "t"))
.await
.unwrap();
sink.write(&record(tracing::Level::INFO, "t"))
.await
.unwrap();
sink.write(&record(tracing::Level::INFO, "t"))
.await
.unwrap();
assert_eq!(
inner.messages.lock().len(),
2,
"budget=2 → exactly 2 writes"
);
assert_eq!(
metrics.logs_dropped(),
1,
"rejected record must count as dropped"
);
}
#[tokio::test]
async fn test_rate_limited_sink_reports_failure_for_budget_refund() {
let inner = Arc::new(CollectingSink {
messages: Mutex::new(Vec::new()),
fail_once: std::sync::atomic::AtomicBool::new(true),
});
let limiter = Arc::new(TokenBucketRateLimit::new(2));
let sink = RateLimitedSink::new(inner.clone() as Arc<dyn LogSink>, limiter.clone());
let r = sink.write(&record(tracing::Level::INFO, "t")).await;
assert!(r.is_err(), "inner failure must propagate");
assert_eq!(
limiter.available(),
2,
"failed write must refund the budget"
);
sink.write(&record(tracing::Level::INFO, "t"))
.await
.unwrap();
assert_eq!(limiter.available(), 1);
assert_eq!(inner.messages.lock().len(), 1);
}
#[test]
fn test_port_is_object_safe() {
fn accepts(_limiter: Arc<dyn SinkRateLimit>) {}
accepts(Arc::new(NoOpRateLimit));
accepts(Arc::new(TokenBucketRateLimit::new(1)));
}
}