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    #[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
54/// Generic asynchronous batch writer with threshold, delayed, and overflow flushing.
55pub 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    /// Creates a writer from product-owned persistence callbacks.
115    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    /// Records one item, scheduling a threshold or delayed flush as needed.
140    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    /// Flushes the current buffer and schedules any remaining buffered items.
177    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    /// Cancels delayed flush tasks. Call [`BufferedBatchWriter::flush`] afterwards during shutdown.
188    pub fn cancel(&self) {
189        self.shutdown_token.cancel();
190    }
191
192    /// Holds the flush lock for deterministic downstream tests.
193    #[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}