1use crate::{ProgressEntry, ProgressListener, Pusher};
4use bytes::{Bytes, BytesMut};
5use std::collections::{BTreeMap, btree_map::Entry};
6
7#[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 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 #[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 #[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 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 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 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 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 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 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 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 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 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 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 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); assert_eq!(pushes[1].0, 2..30); assert_eq!(pushes[1].1.len(), 28);
393 }
394}