1use std::future::Future;
9use std::pin::Pin;
10use std::sync::{
11 Arc, Mutex, MutexGuard as StdMutexGuard,
12 atomic::{AtomicBool, Ordering},
13};
14use std::time::Duration;
15
16use tokio::sync::{Mutex as AsyncMutex, MutexGuard};
17use tokio_util::sync::CancellationToken;
18
19type WriteFuture = Pin<Box<dyn Future<Output = ()> + Send>>;
20type WriteBatch<T> = dyn Fn(Vec<T>) -> WriteFuture + Send + Sync;
21type WriteOne<T> = dyn Fn(T) -> WriteFuture + Send + Sync;
22
23#[derive(Debug, Clone)]
25pub struct BufferedBatchConfig {
26 pub queue_capacity: usize,
28 pub batch_size: usize,
30 pub delayed_flush_after: Duration,
32 pub overflow_label: &'static str,
34}
35
36impl BufferedBatchConfig {
37 pub fn new(
39 queue_capacity: usize,
40 batch_size: usize,
41 delayed_flush_after: Duration,
42 overflow_label: &'static str,
43 ) -> Self {
44 Self {
45 queue_capacity,
46 batch_size,
47 delayed_flush_after,
48 overflow_label,
49 }
50 }
51}
52
53pub struct BufferedBatchWriter<T> {
55 config: BufferedBatchConfig,
56 buffer: Mutex<Vec<T>>,
57 flush_lock: AsyncMutex<()>,
58 flush_pending: AtomicBool,
59 delayed_flush_pending: AtomicBool,
60 shutdown_token: CancellationToken,
61 write_batch: Arc<WriteBatch<T>>,
62 write_one: Arc<WriteOne<T>>,
63}
64
65struct FlushPendingReset<T> {
66 writer: Arc<BufferedBatchWriter<T>>,
67 armed: bool,
68}
69
70impl<T> Drop for FlushPendingReset<T> {
71 fn drop(&mut self) {
72 if self.armed {
73 self.writer.flush_pending.store(false, Ordering::Release);
74 }
75 }
76}
77
78impl<T> FlushPendingReset<T> {
79 fn reset(&mut self) {
80 self.writer.flush_pending.store(false, Ordering::Release);
81 self.armed = false;
82 }
83}
84
85struct DelayedFlushPendingReset<T> {
86 writer: Arc<BufferedBatchWriter<T>>,
87 armed: bool,
88}
89
90impl<T> Drop for DelayedFlushPendingReset<T> {
91 fn drop(&mut self) {
92 if self.armed {
93 self.writer
94 .delayed_flush_pending
95 .store(false, Ordering::Release);
96 }
97 }
98}
99
100impl<T> DelayedFlushPendingReset<T> {
101 fn reset(&mut self) {
102 self.writer
103 .delayed_flush_pending
104 .store(false, Ordering::Release);
105 self.armed = false;
106 }
107}
108
109impl<T> BufferedBatchWriter<T>
110where
111 T: Send + 'static,
112{
113 pub fn new<BatchFn, BatchFuture, OneFn, OneFuture>(
115 config: BufferedBatchConfig,
116 write_batch: BatchFn,
117 write_one: OneFn,
118 ) -> Self
119 where
120 BatchFn: Fn(Vec<T>) -> BatchFuture + Send + Sync + 'static,
121 BatchFuture: Future<Output = ()> + Send + 'static,
122 OneFn: Fn(T) -> OneFuture + Send + Sync + 'static,
123 OneFuture: Future<Output = ()> + Send + 'static,
124 {
125 let batch_capacity = config.batch_size.max(1);
126 Self {
127 config,
128 buffer: Mutex::new(Vec::with_capacity(batch_capacity)),
129 flush_lock: AsyncMutex::new(()),
130 flush_pending: AtomicBool::new(false),
131 delayed_flush_pending: AtomicBool::new(false),
132 shutdown_token: CancellationToken::new(),
133 write_batch: Arc::new(move |items| Box::pin(write_batch(items))),
134 write_one: Arc::new(move |item| Box::pin(write_one(item))),
135 }
136 }
137
138 pub async fn record(self: &Arc<Self>, item: T) {
140 let mut overflow_item = None;
141 let should_flush;
142 let should_schedule_delayed_flush;
143 {
144 let mut buffer = self.lock_buffer();
145 if buffer.len() >= self.config.queue_capacity {
146 overflow_item = Some(item);
147 should_flush = false;
148 should_schedule_delayed_flush = false;
149 } else {
150 let was_empty = buffer.is_empty();
151 buffer.push(item);
152 should_flush = buffer.len() >= self.config.batch_size;
153 should_schedule_delayed_flush = !should_flush && was_empty;
154 }
155 }
156
157 if let Some(item) = overflow_item {
158 tracing::warn!(
159 capacity = self.config.queue_capacity,
160 queue = self.config.overflow_label,
161 "buffered writer queue is full; falling back to direct write"
162 );
163 self.schedule_flush();
164 (self.write_one)(item).await;
165 return;
166 }
167
168 if should_flush {
169 self.schedule_flush();
170 } else if should_schedule_delayed_flush {
171 self.schedule_delayed_flush();
172 }
173 }
174
175 pub async fn flush(self: &Arc<Self>) {
177 let _guard = self.flush_lock.lock().await;
178 self.flush_buffer().await;
179 if self.lock_buffer().is_empty() {
180 self.flush_pending.store(false, Ordering::Release);
181 self.delayed_flush_pending.store(false, Ordering::Release);
182 }
183 self.schedule_buffered_flush();
184 }
185
186 pub fn cancel(&self) {
188 self.shutdown_token.cancel();
189 }
190
191 #[doc(hidden)]
193 pub async fn lock_flush_for_test(&self) -> MutexGuard<'_, ()> {
194 self.flush_lock.lock().await
195 }
196
197 fn schedule_flush(self: &Arc<Self>) {
198 if self
199 .flush_pending
200 .compare_exchange(false, true, Ordering::SeqCst, Ordering::Relaxed)
201 .is_err()
202 {
203 return;
204 }
205
206 let writer = Arc::clone(self);
207 drop(tokio::spawn(async move {
208 let mut pending_reset = FlushPendingReset {
209 writer: Arc::clone(&writer),
210 armed: true,
211 };
212 {
213 let _guard = writer.flush_lock.lock().await;
214 writer.flush_buffer().await;
215 }
216 pending_reset.reset();
217 writer.schedule_buffered_flush();
218 }));
219 }
220
221 fn schedule_delayed_flush(self: &Arc<Self>) {
222 if self
223 .delayed_flush_pending
224 .compare_exchange(false, true, Ordering::SeqCst, Ordering::Relaxed)
225 .is_err()
226 {
227 return;
228 }
229
230 let writer = Arc::clone(self);
231 drop(tokio::spawn(async move {
232 let mut pending_reset = DelayedFlushPendingReset {
233 writer: Arc::clone(&writer),
234 armed: true,
235 };
236 let delayed_flush_after = writer.config.delayed_flush_after;
237 tokio::select! {
238 biased;
239 _ = writer.shutdown_token.cancelled() => return,
240 _ = tokio::time::sleep(delayed_flush_after) => {}
241 }
242
243 {
244 let _guard = writer.flush_lock.lock().await;
245 writer.flush_buffer().await;
246 }
247 pending_reset.reset();
248 writer.schedule_buffered_flush();
249 }));
250 }
251
252 fn schedule_buffered_flush(self: &Arc<Self>) {
253 let buffered_count = self.lock_buffer().len();
254 if buffered_count >= self.config.batch_size {
255 self.schedule_flush();
256 } else if buffered_count > 0 {
257 self.schedule_delayed_flush();
258 }
259 }
260
261 async fn flush_buffer(&self) {
262 let mut items = {
263 let mut buffer = self.lock_buffer();
264 if buffer.is_empty() {
265 return;
266 }
267 std::mem::take(&mut *buffer)
268 };
269
270 let batch_size = self.config.batch_size.max(1);
271 while !items.is_empty() {
272 let chunk_len = items.len().min(batch_size);
273 let chunk = items.drain(..chunk_len).collect::<Vec<_>>();
274 (self.write_batch)(chunk).await;
275 }
276 }
277
278 fn lock_buffer(&self) -> StdMutexGuard<'_, Vec<T>> {
279 match self.buffer.lock() {
280 Ok(guard) => guard,
281 Err(poisoned) => {
282 tracing::warn!(
283 queue = self.config.overflow_label,
284 "buffered writer mutex was poisoned; continuing with recovered buffer"
285 );
286 poisoned.into_inner()
287 }
288 }
289 }
290}
291
292#[cfg(test)]
293mod tests {
294 use std::sync::Arc;
295 use std::sync::atomic::{AtomicUsize, Ordering};
296 use std::time::{Duration, Instant};
297
298 use super::{BufferedBatchConfig, BufferedBatchWriter};
299
300 fn config(delayed_flush_after: Duration) -> BufferedBatchConfig {
301 BufferedBatchConfig::new(5, 3, delayed_flush_after, "test")
302 }
303
304 async fn wait_for_count(count: &AtomicUsize, expected: usize) {
305 let deadline = Instant::now() + Duration::from_secs(2);
306 loop {
307 let current = count.load(Ordering::SeqCst);
308 if current == expected {
309 return;
310 }
311 assert!(
312 current < expected,
313 "count exceeded expected value: expected {expected}, got {current}"
314 );
315 assert!(
316 Instant::now() < deadline,
317 "timed out waiting for count {expected}; last count was {current}"
318 );
319 tokio::time::sleep(Duration::from_millis(10)).await;
320 }
321 }
322
323 #[tokio::test]
324 async fn threshold_batch_flushes_immediately() {
325 let count = Arc::new(AtomicUsize::new(0));
326 let batch_count = Arc::clone(&count);
327 let one_count = Arc::clone(&count);
328 let writer = Arc::new(BufferedBatchWriter::new(
329 config(Duration::from_secs(5)),
330 move |items: Vec<usize>| {
331 let batch_count = Arc::clone(&batch_count);
332 async move {
333 batch_count.fetch_add(items.len(), Ordering::SeqCst);
334 }
335 },
336 move |_item| {
337 let one_count = Arc::clone(&one_count);
338 async move {
339 one_count.fetch_add(1, Ordering::SeqCst);
340 }
341 },
342 ));
343
344 writer.record(1).await;
345 writer.record(2).await;
346 writer.record(3).await;
347
348 wait_for_count(&count, 3).await;
349 writer.cancel();
350 }
351
352 #[tokio::test]
353 async fn partial_batch_flushes_after_delay() {
354 let count = Arc::new(AtomicUsize::new(0));
355 let batch_count = Arc::clone(&count);
356 let one_count = Arc::clone(&count);
357 let writer = Arc::new(BufferedBatchWriter::new(
358 config(Duration::from_millis(20)),
359 move |items: Vec<usize>| {
360 let batch_count = Arc::clone(&batch_count);
361 async move {
362 batch_count.fetch_add(items.len(), Ordering::SeqCst);
363 }
364 },
365 move |_item| {
366 let one_count = Arc::clone(&one_count);
367 async move {
368 one_count.fetch_add(1, Ordering::SeqCst);
369 }
370 },
371 ));
372
373 writer.record(1).await;
374
375 wait_for_count(&count, 1).await;
376 writer.cancel();
377 }
378
379 #[tokio::test]
380 async fn cancel_stops_delayed_flush_until_manual_flush() {
381 let count = Arc::new(AtomicUsize::new(0));
382 let batch_count = Arc::clone(&count);
383 let one_count = Arc::clone(&count);
384 let writer = Arc::new(BufferedBatchWriter::new(
385 config(Duration::from_millis(20)),
386 move |items: Vec<usize>| {
387 let batch_count = Arc::clone(&batch_count);
388 async move {
389 batch_count.fetch_add(items.len(), Ordering::SeqCst);
390 }
391 },
392 move |_item| {
393 let one_count = Arc::clone(&one_count);
394 async move {
395 one_count.fetch_add(1, Ordering::SeqCst);
396 }
397 },
398 ));
399
400 writer.record(1).await;
401 writer.cancel();
402 tokio::time::sleep(Duration::from_millis(60)).await;
403 assert_eq!(count.load(Ordering::SeqCst), 0);
404
405 writer.flush().await;
406 assert_eq!(count.load(Ordering::SeqCst), 1);
407 }
408
409 #[tokio::test]
410 async fn overflow_writes_extra_item_directly_and_flushes_buffer() {
411 let count = Arc::new(AtomicUsize::new(0));
412 let batch_count = Arc::clone(&count);
413 let one_count = Arc::clone(&count);
414 let writer = Arc::new(BufferedBatchWriter::new(
415 config(Duration::from_secs(5)),
416 move |items: Vec<usize>| {
417 let batch_count = Arc::clone(&batch_count);
418 async move {
419 batch_count.fetch_add(items.len(), Ordering::SeqCst);
420 }
421 },
422 move |_item| {
423 let one_count = Arc::clone(&one_count);
424 async move {
425 one_count.fetch_add(1, Ordering::SeqCst);
426 }
427 },
428 ));
429 let flush_guard = writer.lock_flush_for_test().await;
430
431 for index in 0..5 {
432 writer.record(index).await;
433 }
434 writer.record(10_000).await;
435
436 assert_eq!(count.load(Ordering::SeqCst), 1);
437 drop(flush_guard);
438
439 wait_for_count(&count, 6).await;
440 writer.cancel();
441 }
442}