index/bloom_filter/creator/
finalize_segment.rs1use std::pin::Pin;
16use std::sync::Arc;
17use std::sync::atomic::{AtomicUsize, Ordering};
18
19use asynchronous_codec::{FramedRead, FramedWrite};
20use fastbloom::BloomFilter;
21use futures::stream::StreamExt;
22use futures::{AsyncWriteExt, Stream, stream};
23use snafu::ResultExt;
24
25use crate::bloom_filter::PrehashedBuildHasher;
26use crate::bloom_filter::creator::intermediate_codec::IntermediateBloomFilterCodecV1;
27use crate::bloom_filter::error::{IntermediateSnafu, IoSnafu, Result};
28use crate::external_provider::ExternalTempFileProvider;
29
30const MIN_MEMORY_USAGE_THRESHOLD: usize = 1024 * 1024; pub struct FinalizedBloomFilterStorage {
35 false_positive_rate: f64,
37
38 segment_indices: Vec<usize>,
40
41 in_memory: Vec<FinalizedBloomFilterSegment>,
43
44 intermediate_file_id_counter: usize,
46
47 intermediate_prefix: String,
49
50 intermediate_provider: Arc<dyn ExternalTempFileProvider>,
52
53 memory_usage: usize,
55
56 global_memory_usage: Arc<AtomicUsize>,
59
60 global_memory_usage_threshold: Option<usize>,
62
63 flushed_seg_count: usize,
65}
66
67impl FinalizedBloomFilterStorage {
68 pub fn new(
70 false_positive_rate: f64,
71 intermediate_provider: Arc<dyn ExternalTempFileProvider>,
72 global_memory_usage: Arc<AtomicUsize>,
73 global_memory_usage_threshold: Option<usize>,
74 ) -> Self {
75 let external_prefix = format!("intm-bloom-filters-{}", uuid::Uuid::new_v4());
76 Self {
77 false_positive_rate,
78 segment_indices: Vec::new(),
79 in_memory: Vec::new(),
80 intermediate_file_id_counter: 0,
81 intermediate_prefix: external_prefix,
82 intermediate_provider,
83 memory_usage: 0,
84 global_memory_usage,
85 global_memory_usage_threshold,
86 flushed_seg_count: 0,
87 }
88 }
89
90 pub fn memory_usage(&self) -> usize {
92 self.memory_usage
93 }
94
95 pub async fn add(
99 &mut self,
100 elem_hashes: impl IntoIterator<Item = u64>,
101 element_count: usize,
102 ) -> Result<()> {
103 let mut bf = BloomFilter::with_false_pos(self.false_positive_rate)
104 .hasher(PrehashedBuildHasher::default())
105 .expected_items(element_count);
106 for hash in elem_hashes.into_iter() {
107 bf.insert(&hash);
108 }
109
110 let fbf = FinalizedBloomFilterSegment::from(bf, element_count);
111
112 if self.in_memory.last() == Some(&fbf) {
114 self.segment_indices
115 .push(self.flushed_seg_count + self.in_memory.len() - 1);
116 return Ok(());
117 }
118
119 let memory_diff = fbf.bloom_filter_bytes.len();
121 self.memory_usage += memory_diff;
122 self.global_memory_usage
123 .fetch_add(memory_diff, Ordering::Relaxed);
124
125 self.in_memory.push(fbf);
127 self.segment_indices
128 .push(self.flushed_seg_count + self.in_memory.len() - 1);
129
130 if self.memory_usage < MIN_MEMORY_USAGE_THRESHOLD {
134 return Ok(());
135 }
136
137 if let Some(threshold) = self.global_memory_usage_threshold {
139 let global = self.global_memory_usage.load(Ordering::Relaxed);
140
141 if global > threshold {
142 self.flush_in_memory_to_disk().await?;
143
144 self.global_memory_usage
145 .fetch_sub(self.memory_usage, Ordering::Relaxed);
146 self.memory_usage = 0;
147 }
148 }
149
150 Ok(())
151 }
152
153 pub async fn drain(
155 &mut self,
156 ) -> Result<(
157 Vec<usize>,
158 Pin<Box<dyn Stream<Item = Result<FinalizedBloomFilterSegment>> + Send + '_>>,
159 )> {
160 if self.intermediate_file_id_counter == 0 {
162 return Ok((
163 std::mem::take(&mut self.segment_indices),
164 Box::pin(stream::iter(self.in_memory.drain(..).map(Ok))),
165 ));
166 }
167
168 let mut on_disk = self
170 .intermediate_provider
171 .read_all(&self.intermediate_prefix)
172 .await
173 .context(IntermediateSnafu)?;
174 on_disk.sort_unstable_by(|x, y| x.0.cmp(&y.0));
175
176 let streams = on_disk
177 .into_iter()
178 .map(|(_, reader)| FramedRead::new(reader, IntermediateBloomFilterCodecV1::default()));
179
180 let in_memory_stream = stream::iter(self.in_memory.drain(..)).map(Ok);
181 Ok((
182 std::mem::take(&mut self.segment_indices),
183 Box::pin(stream::iter(streams).flatten().chain(in_memory_stream)),
184 ))
185 }
186
187 async fn flush_in_memory_to_disk(&mut self) -> Result<()> {
189 let file_id = self.intermediate_file_id_counter;
190 self.intermediate_file_id_counter += 1;
191 self.flushed_seg_count += self.in_memory.len();
192
193 let file_id = format!("{:08}", file_id);
194 let mut writer = self
195 .intermediate_provider
196 .create(&self.intermediate_prefix, &file_id)
197 .await
198 .context(IntermediateSnafu)?;
199
200 let fw = FramedWrite::new(&mut writer, IntermediateBloomFilterCodecV1::default());
201 if let Err(e) = stream::iter(self.in_memory.drain(..).map(Ok))
203 .forward(fw)
204 .await
205 {
206 writer.close().await.context(IoSnafu)?;
207 writer.flush().await.context(IoSnafu)?;
208 return Err(e);
209 }
210
211 Ok(())
212 }
213}
214
215impl Drop for FinalizedBloomFilterStorage {
216 fn drop(&mut self) {
217 self.global_memory_usage
218 .fetch_sub(self.memory_usage, Ordering::Relaxed);
219 }
220}
221
222#[derive(Debug, Clone, PartialEq, Eq)]
224pub struct FinalizedBloomFilterSegment {
225 pub bloom_filter_bytes: Vec<u8>,
227
228 pub element_count: usize,
230}
231
232impl FinalizedBloomFilterSegment {
233 fn from<S: std::hash::BuildHasher>(bf: BloomFilter<512, S>, elem_count: usize) -> Self {
234 let bf_slice = bf.as_slice();
235 let mut bloom_filter_bytes = Vec::with_capacity(std::mem::size_of_val(bf_slice));
236 for &x in bf_slice {
237 bloom_filter_bytes.extend_from_slice(&x.to_le_bytes());
238 }
239
240 Self {
241 bloom_filter_bytes,
242 element_count: elem_count,
243 }
244 }
245}
246
247#[cfg(test)]
248mod tests {
249 use std::collections::HashMap;
250 use std::sync::Mutex;
251
252 use futures::AsyncRead;
253 use tokio::io::duplex;
254 use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt};
255
256 use super::*;
257 use crate::bloom_filter::creator::tests::u64_vec_from_bytes;
258 use crate::bloom_filter::{SEED, element_hash};
259 use crate::external_provider::MockExternalTempFileProvider;
260
261 #[tokio::test]
262 async fn test_finalized_bloom_filter_storage() {
263 let mut mock_provider = MockExternalTempFileProvider::new();
264
265 let mock_files: Arc<Mutex<HashMap<String, Box<dyn AsyncRead + Unpin + Send>>>> =
266 Arc::new(Mutex::new(HashMap::new()));
267
268 mock_provider.expect_create().returning({
269 let files = Arc::clone(&mock_files);
270 move |file_group, file_id| {
271 assert!(file_group.starts_with("intm-bloom-filters-"));
272 let mut files = files.lock().unwrap();
273 let (writer, reader) = duplex(2 * 1024 * 1024);
274 files.insert(file_id.to_string(), Box::new(reader.compat()));
275 Ok(Box::new(writer.compat_write()))
276 }
277 });
278
279 mock_provider.expect_read_all().returning({
280 let files = Arc::clone(&mock_files);
281 move |file_group| {
282 assert!(file_group.starts_with("intm-bloom-filters-"));
283 let mut files = files.lock().unwrap();
284 Ok(files.drain().collect::<Vec<_>>())
285 }
286 });
287
288 let global_memory_usage = Arc::new(AtomicUsize::new(0));
289 let global_memory_usage_threshold = Some(1024 * 1024); let provider = Arc::new(mock_provider);
291 let mut storage = FinalizedBloomFilterStorage::new(
292 0.01,
293 provider,
294 global_memory_usage.clone(),
295 global_memory_usage_threshold,
296 );
297
298 let elem_count = 2000;
299 let batch = 1000;
300 let dup_batch = 200;
301
302 for i in 0..(batch - dup_batch) {
303 let elems = (elem_count * i..elem_count * (i + 1))
304 .map(|x| element_hash(x.to_string().as_bytes()));
305 storage.add(elems, elem_count).await.unwrap();
306 }
307 for _ in 0..dup_batch {
308 storage.add(Some(element_hash(&[])), 1).await.unwrap();
309 }
310
311 assert!(storage.intermediate_file_id_counter > 0);
313
314 let (indices, mut stream) = storage.drain().await.unwrap();
316 assert_eq!(indices.len(), batch);
317
318 for (i, idx) in indices.iter().enumerate().take(batch - dup_batch) {
319 let segment = stream.next().await.unwrap().unwrap();
320 assert_eq!(segment.element_count, elem_count);
321
322 let v = u64_vec_from_bytes(&segment.bloom_filter_bytes);
323
324 let bf = BloomFilter::from_vec(v)
326 .seed(&SEED)
327 .expected_items(segment.element_count);
328 for elem in (elem_count * i..elem_count * (i + 1)).map(|x| x.to_string().into_bytes()) {
329 assert!(bf.contains(&elem));
330 }
331 assert_eq!(indices[i], *idx);
332 }
333
334 let dup_seg = stream.next().await.unwrap().unwrap();
336 assert_eq!(dup_seg.element_count, 1);
337 assert!(stream.next().await.is_none());
338 assert!(
339 indices[(batch - dup_batch)..batch]
340 .iter()
341 .all(|&x| x == batch - dup_batch)
342 );
343 }
344
345 #[tokio::test]
346 async fn test_finalized_bloom_filter_storage_all_dup() {
347 let mock_provider = MockExternalTempFileProvider::new();
348 let global_memory_usage = Arc::new(AtomicUsize::new(0));
349 let global_memory_usage_threshold = Some(1024 * 1024); let provider = Arc::new(mock_provider);
351 let mut storage = FinalizedBloomFilterStorage::new(
352 0.01,
353 provider,
354 global_memory_usage.clone(),
355 global_memory_usage_threshold,
356 );
357
358 let batch = 1000;
359 for _ in 0..batch {
360 storage.add(Some(element_hash(&[])), 1).await.unwrap();
361 }
362
363 let (indices, mut stream) = storage.drain().await.unwrap();
365
366 let bf = stream.next().await.unwrap().unwrap();
367 assert_eq!(bf.element_count, 1);
368
369 assert!(stream.next().await.is_none());
370
371 assert_eq!(indices.len(), batch);
372 assert!(indices.iter().all(|&x| x == 0));
373 }
374}