Skip to main content

fast_pull/cache/
merge.rs

1//! Pusher cache that merges each flush run into a single contiguous buffer.
2
3use crate::{ProgressEntry, ProgressListener, Pusher};
4use bytes::{Bytes, BytesMut};
5use std::collections::{BTreeMap, btree_map::Entry};
6
7/// Pusher wrapper that buffers chunks and merges each flush run into a single [`Bytes`].
8///
9/// Out-of-order chunks are stored in a `BTreeMap`. When a contiguous run reaches
10/// the high watermark, all chunks in that run are coalesced into one contiguous
11/// [`BytesMut`] before being pushed. This minimizes write calls to the inner pusher
12/// at the cost of an extra memory copy.
13#[derive(Debug)]
14pub struct CacheMergePusher<P> {
15    inner: P,
16    cache: BTreeMap<u64, Bytes>,
17    cache_size: usize,
18    high_watermark: usize,
19    low_watermark: usize,
20}
21
22impl<P: Pusher> CacheMergePusher<P> {
23    /// Wrap `inner` with the given `high_watermark` / `low_watermark` (in bytes).
24    ///
25    /// Eviction to the inner pusher triggers once the buffered size reaches
26    /// `high_watermark`, and stops once it falls back to `low_watermark`.
27    ///
28    /// `low_watermark` must not exceed `high_watermark`. A larger `low_watermark`
29    /// makes `high_watermark` irrelevant, because eviction does nothing until the
30    /// buffered size passes `low_watermark`, which then acts as the sole watermark.
31    ///
32    /// With `high_watermark == low_watermark`, a push that lands the buffered size
33    /// exactly on the watermark evicts nothing; the next push takes it above and
34    /// eviction proceeds as usual.
35    pub const fn new(inner: P, high_watermark: usize, low_watermark: usize) -> Self {
36        Self {
37            inner,
38            cache: BTreeMap::new(),
39            cache_size: 0,
40            high_watermark,
41            low_watermark,
42        }
43    }
44
45    fn evict_until(&mut self, target_size: usize) -> Result<(), P::Error> {
46        if self.cache_size <= target_size {
47            return Ok(());
48        }
49
50        let mut runs: Vec<(u64, usize)> = Vec::with_capacity(self.cache.len());
51        let mut curr_start = None;
52        let mut curr_len = 0;
53        let mut expected_next = 0;
54
55        for (&start, bytes) in &self.cache {
56            let len = bytes.len();
57            if let Some(c_start) = curr_start {
58                if start == expected_next {
59                    curr_len += len;
60                    expected_next += len as u64;
61                } else {
62                    runs.push((c_start, curr_len));
63                    curr_start = Some(start);
64                    curr_len = len;
65                    expected_next = start + len as u64;
66                }
67            } else {
68                curr_start = Some(start);
69                curr_len = len;
70                expected_next = start + len as u64;
71            }
72        }
73        if let Some(c_start) = curr_start {
74            runs.push((c_start, curr_len));
75        }
76        runs.sort_unstable_by_key(|&(_, len)| std::cmp::Reverse(len));
77
78        let mut curr_buf = BytesMut::with_capacity(self.cache_size);
79        let mut err = None;
80        for (start, total_len) in runs {
81            let need_push = err.is_none() && self.cache_size > target_size;
82            let first_bytes = match self.cache.entry(start) {
83                Entry::Occupied(entry) => {
84                    let is_merged = entry.get().len() == total_len;
85                    if !need_push && is_merged {
86                        continue;
87                    }
88                    entry.remove()
89                }
90                Entry::Vacant(_) => unreachable!(),
91            };
92            let chunk = if first_bytes.len() == total_len {
93                first_bytes
94            } else {
95                curr_buf.extend_from_slice(&first_bytes);
96                let mut curr_key = start + first_bytes.len() as u64;
97                let end = start + total_len as u64;
98                while curr_key < end {
99                    let bytes = self.cache.remove(&curr_key).unwrap();
100                    curr_buf.extend_from_slice(&bytes);
101                    curr_key += bytes.len() as u64;
102                }
103                curr_buf.split().freeze()
104            };
105            if need_push {
106                let end = start + total_len as u64;
107                let range = start..end;
108                self.cache_size -= total_len;
109                if let Err((e, ret_bytes)) = self.inner.push(&range, chunk) {
110                    err = Some(e);
111                    if !ret_bytes.is_empty() {
112                        let written = total_len.saturating_sub(ret_bytes.len());
113                        self.cache_size += ret_bytes.len();
114                        if let Some(old) = self.cache.insert(start + written as u64, ret_bytes) {
115                            self.cache_size -= old.len();
116                        }
117                    }
118                }
119            } else if let Some(old) = self.cache.insert(start, chunk) {
120                self.cache_size -= old.len();
121            }
122        }
123        err.map_or(Ok(()), Err)
124    }
125}
126
127impl<P: Pusher> Pusher for CacheMergePusher<P> {
128    type Error = P::Error;
129
130    fn set_listener(&mut self, cb: ProgressListener) {
131        self.inner.set_listener(cb);
132    }
133
134    fn push(&mut self, range: &ProgressEntry, bytes: Bytes) -> Result<(), (Self::Error, Bytes)> {
135        if bytes.is_empty() {
136            return Ok(());
137        }
138
139        self.cache_size += bytes.len();
140        if let Some(old_bytes) = self.cache.insert(range.start, bytes) {
141            self.cache_size -= old_bytes.len();
142        }
143
144        if self.cache_size >= self.high_watermark
145            && let Err(e) = self.evict_until(self.low_watermark)
146        {
147            return Err((e, Bytes::new()));
148        }
149
150        Ok(())
151    }
152
153    fn flush(&mut self) -> Result<(), Self::Error> {
154        self.evict_until(0)?;
155        self.inner.flush()
156    }
157}
158
159#[cfg(test)]
160mod tests {
161    #![allow(clippy::unwrap_used)]
162    use super::*;
163    use std::sync::atomic::{AtomicBool, Ordering};
164    use std::sync::{Arc, Mutex};
165
166    /// A `Pusher` that records everything pushed into a shared buffer so tests
167    /// can inspect it without reaching into `CacheMergePusher`'s private fields.
168    #[derive(Clone)]
169    struct SharedSink {
170        pushes: Arc<Mutex<Vec<(ProgressEntry, Bytes)>>>,
171        fail_next: Arc<AtomicBool>,
172        listener_set: Arc<AtomicBool>,
173    }
174    impl SharedSink {
175        fn new() -> Self {
176            Self {
177                pushes: Arc::new(Mutex::new(Vec::new())),
178                fail_next: Arc::new(AtomicBool::new(false)),
179                listener_set: Arc::new(AtomicBool::new(false)),
180            }
181        }
182    }
183    impl Pusher for SharedSink {
184        type Error = std::io::Error;
185        fn set_listener(&mut self, _: ProgressListener) {
186            self.listener_set.store(true, Ordering::SeqCst);
187        }
188        fn push(
189            &mut self,
190            range: &ProgressEntry,
191            bytes: Bytes,
192        ) -> Result<(), (Self::Error, Bytes)> {
193            if self.fail_next.fetch_and(false, Ordering::SeqCst) {
194                return Err((std::io::Error::other("boom"), bytes));
195            }
196            self.pushes.lock().unwrap().push((range.clone(), bytes));
197            Ok(())
198        }
199        fn flush(&mut self) -> Result<(), Self::Error> {
200            Ok(())
201        }
202    }
203
204    /// Inner pusher that writes only the first 2 bytes of the chunk on its first call,
205    /// then fails returning the unwritten tail as `rem`. Exercises `evict_until`'s
206    /// partial-write re-buffering for a *merged* run.
207    #[derive(Clone)]
208    struct PartialSinkMerge {
209        pushes: Arc<Mutex<Vec<(ProgressEntry, Bytes)>>>,
210        partial: Arc<AtomicBool>,
211    }
212    impl Pusher for PartialSinkMerge {
213        type Error = std::io::Error;
214        fn push(
215            &mut self,
216            range: &ProgressEntry,
217            bytes: Bytes,
218        ) -> Result<(), (Self::Error, Bytes)> {
219            self.pushes
220                .lock()
221                .unwrap()
222                .push((range.clone(), bytes.clone()));
223            if !self.partial.swap(true, Ordering::SeqCst) {
224                let rem = bytes.slice(2..);
225                return Err((std::io::Error::other("partial"), rem));
226            }
227            Ok(())
228        }
229    }
230
231    fn bb(s: &str) -> Bytes {
232        Bytes::copy_from_slice(s.as_bytes())
233    }
234
235    #[test]
236    fn test_cache_merge_evicts_contiguous_run() {
237        let sink = SharedSink::new();
238        let mut p = CacheMergePusher::new(sink.clone(), 30, 0);
239        // Out-of-order insertion; not yet at watermark.
240        p.push(&(0..10), bb(&"A".repeat(10))).unwrap();
241        p.push(&(20..30), bb(&"C".repeat(10))).unwrap();
242        assert!(sink.pushes.lock().unwrap().is_empty());
243        // This insertion reaches the high watermark and triggers a merge+evict.
244        p.push(&(10..20), bb(&"B".repeat(10))).unwrap();
245        let pushes = sink.pushes.lock().unwrap();
246        assert_eq!(pushes.len(), 1);
247        assert_eq!(pushes[0].0, 0..30);
248        assert_eq!(pushes[0].1.len(), 30);
249        // Merged bytes preserve ascending order: A(0..10) B(10..20) C(20..30).
250        assert_eq!(&pushes[0].1[..], b"AAAAAAAAAABBBBBBBBBBCCCCCCCCCC");
251        drop(pushes);
252    }
253
254    #[test]
255    fn test_cache_merge_no_evict_below_watermark() {
256        let sink = SharedSink::new();
257        let mut p = CacheMergePusher::new(sink.clone(), 100, 0);
258        p.push(&(0..10), bb(&"A".repeat(10))).unwrap();
259        p.push(&(10..20), bb(&"B".repeat(10))).unwrap();
260        assert!(sink.pushes.lock().unwrap().is_empty());
261        // Explicit flush must drain the buffered run to the inner pusher.
262        p.flush().unwrap();
263        let pushes = sink.pushes.lock().unwrap();
264        assert_eq!(pushes.len(), 1);
265        assert_eq!(pushes[0].0, 0..20);
266        drop(pushes);
267    }
268
269    #[test]
270    fn test_cache_merge_inner_failure_propagates() {
271        let sink = SharedSink::new();
272        sink.fail_next.store(true, Ordering::SeqCst);
273        let mut p = CacheMergePusher::new(sink, 10, 0);
274        let res = p.push(&(0..10), bb(&"A".repeat(10)));
275        assert!(res.is_err());
276    }
277
278    #[test]
279    fn test_cache_merge_empty_bytes_is_noop() {
280        let sink = SharedSink::new();
281        let mut p = CacheMergePusher::new(sink.clone(), 1, 0);
282        p.push(&(0..0), Bytes::new()).unwrap();
283        assert!(sink.pushes.lock().unwrap().is_empty());
284    }
285
286    #[test]
287    fn test_cache_merge_partial_overwrite_replaces() {
288        // A chunk at an already-cached start position replaces the old bytes.
289        let sink = SharedSink::new();
290        let mut p = CacheMergePusher::new(sink.clone(), 100, 0);
291        p.push(&(0..10), bb(&"A".repeat(10))).unwrap();
292        p.push(&(0..10), bb(&"B".repeat(10))).unwrap();
293        p.flush().unwrap();
294        let pushes = sink.pushes.lock().unwrap();
295        assert_eq!(pushes.len(), 1);
296        assert_eq!(&pushes[0].1[..], b"BBBBBBBBBB");
297        drop(pushes);
298    }
299
300    #[test]
301    fn test_cache_merge_set_listener_forwards() {
302        // Lines 120-122 (forwarding) and 173 (inner set_listener): the listener must
303        // reach the inner sink.
304        let sink = SharedSink::new();
305        let mut p = CacheMergePusher::new(sink.clone(), 100, 0);
306        p.set_listener(Box::new(|_| {}));
307        assert!(sink.listener_set.load(Ordering::SeqCst));
308    }
309
310    #[test]
311    fn test_cache_merge_equal_watermarks_skip_only_the_exact_hit() {
312        // With high == low, `evict_until` returns early only while `cache_size` is
313        // exactly on the watermark. It is not a permanent no-op: the next push takes
314        // the buffer above the watermark and eviction runs normally.
315        let sink = SharedSink::new();
316        let mut p = CacheMergePusher::new(sink.clone(), 10, 10);
317        p.push(&(0..10), bb(&"A".repeat(10))).unwrap();
318        assert!(sink.pushes.lock().unwrap().is_empty());
319
320        p.push(&(10..11), bb("B")).unwrap();
321        let pushes = sink.pushes.lock().unwrap();
322        assert_eq!(pushes.len(), 1);
323        assert_eq!(pushes[0].0, 0..11);
324        assert_eq!(&pushes[0].1[..], format!("{}B", "A".repeat(10)).as_bytes());
325    }
326
327    #[test]
328    fn test_cache_merge_gap_segmentation_and_reinsert() {
329        // Covers 54-58 (gap segmentation), 84-96 (multi-chunk merge) and 109-111
330        // (a multi-chunk run re-inserted when eviction stops at the low watermark).
331        //
332        // Layout at eviction time (high=200, low=70):
333        //   runA [0..100)  two chunks  -> 100 bytes (longest)
334        //   runB [200..260) two chunks -> 60 bytes
335        //   runC [300..340) two chunks -> 40 bytes (multi-chunk, stays buffered)
336        //   runD [400..420) one chunk  -> 20 bytes (single, stays buffered)
337        let sink = SharedSink::new();
338        let mut p = CacheMergePusher::new(sink.clone(), 200, 70);
339        p.push(&(0..50), bb(&"A".repeat(50))).unwrap();
340        p.push(&(50..100), bb(&"A".repeat(50))).unwrap();
341        p.push(&(200..230), bb(&"B".repeat(30))).unwrap();
342        p.push(&(230..260), bb(&"B".repeat(30))).unwrap();
343        p.push(&(300..320), bb(&"C".repeat(20))).unwrap();
344        p.push(&(320..340), bb(&"C".repeat(20))).unwrap();
345        p.push(&(400..420), bb(&"D".repeat(20))).unwrap();
346
347        let pushes = sink.pushes.lock().unwrap();
348        // runA and runB were pushed; runC/runD remain buffered below the low mark.
349        assert_eq!(pushes.len(), 2);
350        assert_eq!(pushes[0].0, 0..100);
351        assert_eq!(pushes[1].0, 200..260);
352        drop(pushes);
353    }
354
355    #[test]
356    fn test_cache_merge_reaches_low_watermark_stops_single() {
357        // Line 78: once cache_size falls to <= low_watermark, a remaining single-chunk
358        // run is `continue`d (left in the cache) rather than pushed.
359        let sink = SharedSink::new();
360        let mut p = CacheMergePusher::new(sink.clone(), 200, 70);
361        p.push(&(0..100), bb(&"A".repeat(100))).unwrap();
362        p.push(&(200..240), bb(&"B".repeat(40))).unwrap();
363        p.push(&(300..340), bb(&"C".repeat(40))).unwrap();
364        p.push(&(400..420), bb(&"D".repeat(20))).unwrap();
365
366        let pushes = sink.pushes.lock().unwrap();
367        assert_eq!(pushes.len(), 2);
368        assert_eq!(pushes[0].0, 0..100);
369        assert_eq!(pushes[1].0, 200..240);
370        drop(pushes);
371    }
372
373    #[test]
374    fn merge_evict_partial_write_rebuffers_merged_tail_at_correct_offset() {
375        // A contiguous run of three chunks [0..30) is merged into one 30-byte chunk and
376        // evicted; a partial inner write (first 2 bytes persisted, tail returned) must
377        // re-buffer the 28-byte tail at offset 2 and retry it on flush.
378        let sink = PartialSinkMerge {
379            pushes: Arc::new(Mutex::new(Vec::new())),
380            partial: Arc::new(AtomicBool::new(false)),
381        };
382        let mut p = CacheMergePusher::new(sink.clone(), 30, 0);
383        p.push(&(0..10), bb(&"A".repeat(10))).unwrap();
384        p.push(&(10..20), bb(&"B".repeat(10))).unwrap();
385        let res = p.push(&(20..30), bb(&"C".repeat(10)));
386        assert!(res.is_err());
387        p.flush().unwrap();
388        let pushes = sink.pushes.lock().unwrap();
389        assert_eq!(pushes.len(), 2);
390        assert_eq!(pushes[0].0, 0..30); // merged run handed to inner
391        assert_eq!(pushes[1].0, 2..30); // tail re-buffered at correct offset
392        assert_eq!(pushes[1].1.len(), 28);
393    }
394}