use crate::backend::CacheBackend;
use crate::error::{OxCacheError, OxCacheResult};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{Mutex, Notify};
#[derive(Debug, Clone)]
struct BatchEntry {
key: Arc<str>,
value: Arc<Vec<u8>>,
ttl: Option<Duration>,
}
pub struct BatchWriter {
backend: Arc<dyn CacheBackend>,
buffer: Mutex<Vec<BatchEntry>>,
capacity: usize,
flush_interval: Duration,
reject_when_full: bool,
stop_notify: Arc<Notify>,
}
impl BatchWriter {
pub fn builder(backend: Arc<dyn CacheBackend>) -> BatchWriterBuilder {
BatchWriterBuilder::new(backend)
}
pub async fn enqueue(
&self,
key: Arc<str>,
value: Arc<Vec<u8>>,
ttl: Option<Duration>,
) -> OxCacheResult<()> {
let should_flush = {
let mut buf = self.buffer.lock().await;
if self.reject_when_full && buf.len() >= self.capacity {
return Err(OxCacheError::BufferFull(format!(
"batch writer buffer at capacity {}",
self.capacity
)));
}
buf.push(BatchEntry { key, value, ttl });
!self.reject_when_full && buf.len() >= self.capacity
};
if should_flush {
self.flush().await?;
}
Ok(())
}
pub async fn flush(&self) -> OxCacheResult<()> {
let entries: Vec<BatchEntry> = {
let mut buf = self.buffer.lock().await;
std::mem::take(&mut *buf)
};
if entries.is_empty() {
return Ok(());
}
for entry in &entries {
self.backend
.set(entry.key.clone(), entry.value.clone(), entry.ttl)
.await?;
}
Ok(())
}
pub async fn pending(&self) -> usize {
self.buffer.lock().await.len()
}
pub fn start(self: &Arc<Self>) {
let writer = Arc::clone(self);
let stop = Arc::clone(&self.stop_notify);
let interval = self.flush_interval;
tokio::spawn(async move {
loop {
tokio::select! {
_ = tokio::time::sleep(interval) => {
let _ = writer.flush().await;
}
_ = stop.notified() => {
let _ = writer.flush().await;
return;
}
}
}
});
}
pub async fn stop(&self) -> OxCacheResult<()> {
self.stop_notify.notify_one();
tokio::time::sleep(Duration::from_millis(10)).await;
self.flush().await
}
}
pub struct BatchWriterBuilder {
backend: Arc<dyn CacheBackend>,
capacity: usize,
flush_interval: Duration,
reject_when_full: bool,
}
impl BatchWriterBuilder {
fn new(backend: Arc<dyn CacheBackend>) -> Self {
Self {
backend,
capacity: 100,
flush_interval: Duration::from_secs(5),
reject_when_full: false,
}
}
pub fn capacity(mut self, cap: usize) -> Self {
self.capacity = cap.max(1);
self
}
pub fn flush_interval(mut self, interval: Duration) -> Self {
self.flush_interval = interval;
self
}
pub fn reject_when_full(mut self, reject: bool) -> Self {
self.reject_when_full = reject;
self
}
pub fn build(self) -> BatchWriter {
BatchWriter {
backend: self.backend,
buffer: Mutex::new(Vec::new()),
capacity: self.capacity,
flush_interval: self.flush_interval,
reject_when_full: self.reject_when_full,
stop_notify: Arc::new(Notify::new()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::MokaMemoryBackend;
fn test_backend() -> Arc<dyn CacheBackend> {
Arc::new(MokaMemoryBackend::builder().capacity(1000).build())
}
#[tokio::test]
async fn batch_writer_enqueues_and_flushes() {
let backend = test_backend();
let bw = Arc::new(
BatchWriter::builder(backend.clone())
.capacity(10)
.flush_interval(Duration::from_secs(60))
.build(),
);
for i in 0..3u8 {
bw.enqueue(Arc::from(format!("k{i}")), Arc::new(vec![i]), None)
.await
.expect("enqueue");
}
assert_eq!(bw.pending().await, 3);
bw.flush().await.expect("flush");
assert_eq!(bw.pending().await, 0);
let val = backend.get("k1").await.expect("get");
assert_eq!(val, Some(vec![1]));
}
#[tokio::test]
async fn batch_writer_auto_flushes_at_capacity() {
let backend = test_backend();
let bw = Arc::new(
BatchWriter::builder(backend.clone())
.capacity(3)
.flush_interval(Duration::from_secs(60))
.build(),
);
for i in 0..3u8 {
bw.enqueue(Arc::from(format!("auto{i}")), Arc::new(vec![i]), None)
.await
.expect("enqueue");
}
assert_eq!(bw.pending().await, 0);
let val = backend.get("auto2").await.expect("get");
assert_eq!(val, Some(vec![2]));
}
#[tokio::test]
async fn batch_writer_stop_drains_remaining() {
let backend = test_backend();
let bw = Arc::new(
BatchWriter::builder(backend.clone())
.capacity(100)
.flush_interval(Duration::from_secs(60))
.build(),
);
bw.start();
bw.enqueue(Arc::from("drain"), Arc::new(b"data".to_vec()), None)
.await
.expect("enqueue");
assert_eq!(bw.pending().await, 1);
bw.stop().await.expect("stop");
let val = backend.get("drain").await.expect("get");
assert_eq!(val, Some(b"data".to_vec()));
}
#[tokio::test]
async fn batch_writer_flush_empty_is_noop() {
let bw = BatchWriter::builder(test_backend()).build();
bw.flush().await.expect("flush empty");
}
#[test]
fn batch_writer_builder_defaults() {
let bw = BatchWriter::builder(test_backend()).build();
assert_eq!(bw.capacity, 100);
assert_eq!(bw.flush_interval, Duration::from_secs(5));
assert!(!bw.reject_when_full);
}
#[tokio::test]
async fn batch_writer_rejects_when_full_with_recoverable_buffer_full() {
let backend = test_backend();
let bw = BatchWriter::builder(backend.clone())
.capacity(2)
.flush_interval(Duration::from_secs(60))
.reject_when_full(true)
.build();
bw.enqueue(Arc::from("full-k1"), Arc::new(b"v1".to_vec()), None)
.await
.expect("第 1 条应入队");
bw.enqueue(Arc::from("full-k2"), Arc::new(b"v2".to_vec()), None)
.await
.expect("第 2 条应入队");
let err = bw
.enqueue(Arc::from("full-k3"), Arc::new(b"v3".to_vec()), None)
.await
.expect_err("缓冲打满应拒绝入队");
assert!(
matches!(err, OxCacheError::BufferFull(_)),
"应为 BufferFull,实际: {err:?}"
);
assert_eq!(err.code(), "OXCACHE_019");
assert!(err.is_recoverable(), "BufferFull 应是可恢复错误");
assert_eq!(bw.pending().await, 2, "被拒绝的条目不应入缓冲");
bw.flush().await.expect("flush");
assert_eq!(
backend.get("full-k1").await.expect("get"),
Some(b"v1".to_vec())
);
bw.enqueue(Arc::from("full-k4"), Arc::new(b"v4".to_vec()), None)
.await
.expect("flush 后应恢复入队");
}
}