Skip to main content

aster_forge_runtime/
buffered.rs

1//! Generic buffered batch writer for runtime side-effect queues.
2//!
3//! This module owns the product-neutral mechanics behind a common Aster pattern: accept records
4//! quickly, flush when a batch threshold is reached, flush a partial batch after a short delay, and
5//! fall back to a direct write when the in-memory queue is full. Products still own the record type,
6//! persistence calls, error mapping, and any audit or metrics semantics.
7
8use 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/// Runtime buffering policy for [`BufferedBatchWriter`].
24#[derive(Debug, Clone)]
25pub struct BufferedBatchConfig {
26    /// Maximum number of records kept in memory before overflow records use the direct writer.
27    pub queue_capacity: usize,
28    /// Number of queued records that triggers an immediate flush.
29    pub batch_size: usize,
30    /// Delay before flushing a partial batch.
31    pub delayed_flush_after: Duration,
32    /// Label used in overflow warning logs.
33    pub overflow_label: &'static str,
34}
35
36impl BufferedBatchConfig {
37    /// Creates a buffering policy.
38    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
53/// Generic asynchronous batch writer with threshold, delayed, and overflow flushing.
54pub 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    /// Creates a writer from product-owned persistence callbacks.
114    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    /// Records one item, scheduling a threshold or delayed flush as needed.
139    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    /// Flushes the current buffer and schedules any remaining buffered items.
176    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    /// Cancels delayed flush tasks. Call [`BufferedBatchWriter::flush`] afterwards during shutdown.
187    pub fn cancel(&self) {
188        self.shutdown_token.cancel();
189    }
190
191    /// Holds the flush lock for deterministic downstream tests.
192    #[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}