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