Skip to main content

alopex_core/sql/
spill.rs

1//! Generic external-sort spill primitives.
2
3use std::cmp::Ordering;
4use std::collections::BinaryHeap;
5use std::fs::{self, File, OpenOptions};
6use std::io::{BufReader, BufWriter, Read, Write};
7use std::path::{Path, PathBuf};
8use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
9use std::sync::Arc;
10
11use crate::sql::stream::MemoryPolicy;
12use crate::{Error, Result};
13
14type KeyDecoder<K> = dyn Fn(&[u8]) -> Result<K>;
15type RowDecoder<T> = dyn Fn(u64, &[u8]) -> Result<T>;
16type KeyComparator<K> = dyn Fn(&K, &K) -> Ordering;
17type RowId<T> = dyn Fn(&T) -> u64;
18
19static SPILL_COUNTER: AtomicU64 = AtomicU64::new(0);
20
21/// Write one sorted spill run to disk and clear the in-memory entries.
22pub fn spill_run<T, K, C, R, EK, ER>(
23    entries: &mut Vec<(T, K)>,
24    policy: &MemoryPolicy,
25    prefix: &str,
26    mut compare_keys: C,
27    row_id: R,
28    encode_key: EK,
29    encode_row: ER,
30) -> Result<PathBuf>
31where
32    C: FnMut(&K, &K) -> Ordering,
33    R: Fn(&T) -> u64,
34    EK: Fn(&K) -> Vec<u8>,
35    ER: Fn(&T) -> Vec<u8>,
36{
37    let directory = policy.spill_directory().ok_or_else(|| Error::SpillFailed {
38        reason: "sort spill: spill directory not configured".into(),
39    })?;
40    ensure_spill_dir(directory)?;
41    let (path, file) = create_spill_file(directory, prefix)?;
42    let mut writer = BufWriter::new(file);
43
44    entries.sort_by(|a, b| compare_keys(&a.1, &b.1));
45
46    let mut bytes_written = 0u64;
47    for (row, keys) in entries.iter() {
48        let key_bytes = encode_key(keys);
49        let row_bytes = encode_row(row);
50        let key_len = u32::try_from(key_bytes.len()).map_err(|_| Error::SpillFailed {
51            reason: "sort spill: sort key size exceeds u32::MAX".into(),
52        })?;
53        let row_len = u32::try_from(row_bytes.len()).map_err(|_| Error::SpillFailed {
54            reason: "sort spill: row size exceeds u32::MAX".into(),
55        })?;
56
57        writer
58            .write_all(&row_id(row).to_le_bytes())
59            .map_err(|err| spill_io_error("sort spill", err))?;
60        writer
61            .write_all(&key_len.to_le_bytes())
62            .map_err(|err| spill_io_error("sort spill", err))?;
63        writer
64            .write_all(&row_len.to_le_bytes())
65            .map_err(|err| spill_io_error("sort spill", err))?;
66        writer
67            .write_all(&key_bytes)
68            .map_err(|err| spill_io_error("sort spill", err))?;
69        writer
70            .write_all(&row_bytes)
71            .map_err(|err| spill_io_error("sort spill", err))?;
72        bytes_written = bytes_written
73            .saturating_add(8)
74            .saturating_add(4)
75            .saturating_add(4)
76            .saturating_add(key_bytes.len() as u64)
77            .saturating_add(row_bytes.len() as u64);
78    }
79
80    writer
81        .flush()
82        .map_err(|err| spill_io_error("sort spill", err))?;
83    policy.record_spill(bytes_written, 1);
84    entries.clear();
85
86    Ok(path)
87}
88
89/// Ensure the spill directory exists.
90pub fn ensure_spill_dir(directory: &Path) -> Result<()> {
91    fs::create_dir_all(directory).map_err(|err| spill_io_error("sort spill", err))?;
92    Ok(())
93}
94
95/// Create a unique spill file in `directory` using `prefix`.
96pub fn create_spill_file(directory: &Path, prefix: &str) -> Result<(PathBuf, File)> {
97    for _ in 0..16 {
98        let counter = SPILL_COUNTER.fetch_add(1, AtomicOrdering::Relaxed);
99        let timestamp = std::time::SystemTime::now()
100            .duration_since(std::time::UNIX_EPOCH)
101            .unwrap_or_default()
102            .as_nanos();
103        let path = directory.join(format!("{prefix}-{timestamp}-{counter}.bin"));
104        match OpenOptions::new().create_new(true).write(true).open(&path) {
105            Ok(file) => return Ok((path, file)),
106            Err(err) if err.kind() == std::io::ErrorKind::AlreadyExists => continue,
107            Err(err) => return Err(spill_io_error("sort spill", err)),
108        }
109    }
110    Err(Error::SpillFailed {
111        reason: "sort spill: failed to allocate spill file".into(),
112    })
113}
114
115/// Convert a spill I/O failure into the stable core spill error variant.
116pub fn spill_io_error(operation: &str, err: impl std::fmt::Display) -> Error {
117    Error::SpillFailed {
118        reason: format!("{operation}: {err}"),
119    }
120}
121
122/// One decoded entry from a spill run.
123pub struct SpillEntry<T, K> {
124    /// Decoded row payload.
125    pub row: T,
126    /// Decoded sort key payload.
127    pub keys: K,
128}
129
130/// Reader for one spill run file.
131pub struct SpillRunReader<T, K> {
132    path: PathBuf,
133    reader: BufReader<File>,
134    decode_key: Arc<KeyDecoder<K>>,
135    decode_row: Arc<RowDecoder<T>>,
136}
137
138impl<T, K> SpillRunReader<T, K> {
139    /// Open a spill run and configure decoders for key and row payloads.
140    pub fn open<DK, DR>(path: PathBuf, decode_key: DK, decode_row: DR) -> Result<Self>
141    where
142        DK: Fn(&[u8]) -> Result<K> + 'static,
143        DR: Fn(u64, &[u8]) -> Result<T> + 'static,
144    {
145        Self::open_with_decoders(path, Arc::new(decode_key), Arc::new(decode_row))
146    }
147
148    fn open_with_decoders(
149        path: PathBuf,
150        decode_key: Arc<KeyDecoder<K>>,
151        decode_row: Arc<RowDecoder<T>>,
152    ) -> Result<Self> {
153        let file = File::open(&path).map_err(|err| spill_io_error("sort spill", err))?;
154        Ok(Self {
155            path,
156            reader: BufReader::new(file),
157            decode_key,
158            decode_row,
159        })
160    }
161
162    /// Read the next spill entry, or `None` at end of run.
163    pub fn next_entry(&mut self) -> Result<Option<SpillEntry<T, K>>> {
164        let mut row_id_buf = [0u8; 8];
165        match self.reader.read_exact(&mut row_id_buf) {
166            Ok(()) => {}
167            Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
168            Err(err) => return Err(spill_io_error("sort spill", err)),
169        }
170        let row_id = u64::from_le_bytes(row_id_buf);
171        let key_len = self.read_u32()?;
172        let row_len = self.read_u32()?;
173
174        let mut key_bytes = vec![0u8; key_len as usize];
175        self.reader
176            .read_exact(&mut key_bytes)
177            .map_err(|err| spill_io_error("sort spill", err))?;
178        let mut row_bytes = vec![0u8; row_len as usize];
179        self.reader
180            .read_exact(&mut row_bytes)
181            .map_err(|err| spill_io_error("sort spill", err))?;
182
183        let keys = (self.decode_key)(&key_bytes)?;
184        let row = (self.decode_row)(row_id, &row_bytes)?;
185
186        Ok(Some(SpillEntry { row, keys }))
187    }
188
189    fn read_u32(&mut self) -> Result<u32> {
190        let mut buf = [0u8; 4];
191        self.reader
192            .read_exact(&mut buf)
193            .map_err(|err| spill_io_error("sort spill", err))?;
194        Ok(u32::from_le_bytes(buf))
195    }
196}
197
198impl<T, K> Drop for SpillRunReader<T, K> {
199    fn drop(&mut self) {
200        let _ = fs::remove_file(&self.path);
201    }
202}
203
204/// K-way merge iterator over sorted spill run files.
205pub struct SpillMergeIterator<T, K> {
206    compare: Arc<KeyComparator<K>>,
207    row_id: Arc<RowId<T>>,
208    readers: Vec<SpillRunReader<T, K>>,
209    heap: BinaryHeap<SpillHeapItem<T, K>>,
210}
211
212impl<T, K> SpillMergeIterator<T, K> {
213    /// Open spill runs and initialize the merge heap.
214    pub fn new<C, R, DK, DR>(
215        runs: Vec<PathBuf>,
216        compare_keys: C,
217        row_id: R,
218        decode_key: DK,
219        decode_row: DR,
220    ) -> Result<Self>
221    where
222        C: Fn(&K, &K) -> Ordering + 'static,
223        R: Fn(&T) -> u64 + 'static,
224        DK: Fn(&[u8]) -> Result<K> + 'static,
225        DR: Fn(u64, &[u8]) -> Result<T> + 'static,
226    {
227        let compare: Arc<KeyComparator<K>> = Arc::new(compare_keys);
228        let row_id: Arc<RowId<T>> = Arc::new(row_id);
229        let decode_key: Arc<KeyDecoder<K>> = Arc::new(decode_key);
230        let decode_row: Arc<RowDecoder<T>> = Arc::new(decode_row);
231        let mut readers = Vec::with_capacity(runs.len());
232        let mut heap = BinaryHeap::new();
233
234        for (idx, path) in runs.into_iter().enumerate() {
235            let mut reader = SpillRunReader::open_with_decoders(
236                path,
237                Arc::clone(&decode_key),
238                Arc::clone(&decode_row),
239            )?;
240            if let Some(entry) = reader.next_entry()? {
241                heap.push(SpillHeapItem {
242                    run_idx: idx,
243                    row: entry.row,
244                    keys: entry.keys,
245                    compare: Arc::clone(&compare),
246                    row_id: Arc::clone(&row_id),
247                });
248            }
249            readers.push(reader);
250        }
251
252        Ok(Self {
253            compare,
254            row_id,
255            readers,
256            heap,
257        })
258    }
259
260    /// Return the next globally sorted row from the spill runs.
261    pub fn next_item(&mut self) -> Option<Result<T>> {
262        let item = self.heap.pop()?;
263        let row = item.row;
264        let run_idx = item.run_idx;
265
266        match self.readers[run_idx].next_entry() {
267            Ok(Some(entry)) => {
268                self.heap.push(SpillHeapItem {
269                    run_idx,
270                    row: entry.row,
271                    keys: entry.keys,
272                    compare: Arc::clone(&self.compare),
273                    row_id: Arc::clone(&self.row_id),
274                });
275            }
276            Ok(None) => {}
277            Err(err) => return Some(Err(err)),
278        }
279
280        Some(Ok(row))
281    }
282}
283
284impl<T, K> Iterator for SpillMergeIterator<T, K> {
285    type Item = Result<T>;
286
287    fn next(&mut self) -> Option<Self::Item> {
288        self.next_item()
289    }
290}
291
292struct SpillHeapItem<T, K> {
293    run_idx: usize,
294    row: T,
295    keys: K,
296    compare: Arc<KeyComparator<K>>,
297    row_id: Arc<RowId<T>>,
298}
299
300impl<T, K> PartialEq for SpillHeapItem<T, K> {
301    fn eq(&self, other: &Self) -> bool {
302        (self.compare)(&self.keys, &other.keys) == Ordering::Equal
303            && self.run_idx == other.run_idx
304            && (self.row_id)(&self.row) == (self.row_id)(&other.row)
305    }
306}
307
308impl<T, K> Eq for SpillHeapItem<T, K> {}
309
310impl<T, K> PartialOrd for SpillHeapItem<T, K> {
311    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
312        Some(self.cmp(other))
313    }
314}
315
316impl<T, K> Ord for SpillHeapItem<T, K> {
317    fn cmp(&self, other: &Self) -> Ordering {
318        let order = (self.compare)(&self.keys, &other.keys);
319        let order = if order == Ordering::Equal {
320            self.run_idx
321                .cmp(&other.run_idx)
322                .then_with(|| (self.row_id)(&self.row).cmp(&(self.row_id)(&other.row)))
323        } else {
324            order
325        };
326        order.reverse()
327    }
328}
329
330#[cfg(all(test, not(target_arch = "wasm32")))]
331mod tests {
332    use super::{spill_io_error, spill_run, SpillMergeIterator, SpillRunReader};
333    use crate::sql::stream::{MemoryPolicy, SpillPolicy};
334    use crate::Error;
335
336    #[derive(Debug, Clone, PartialEq, Eq)]
337    struct TestRow {
338        row_id: u64,
339        value: i64,
340    }
341
342    fn encode_i64(value: &i64) -> Vec<u8> {
343        value.to_le_bytes().to_vec()
344    }
345
346    fn decode_i64(bytes: &[u8]) -> crate::Result<i64> {
347        let data: [u8; 8] = bytes.try_into().map_err(|_| Error::SpillFailed {
348            reason: "invalid i64 bytes".into(),
349        })?;
350        Ok(i64::from_le_bytes(data))
351    }
352
353    fn encode_row(row: &TestRow) -> Vec<u8> {
354        row.value.to_le_bytes().to_vec()
355    }
356
357    fn decode_row(row_id: u64, bytes: &[u8]) -> crate::Result<TestRow> {
358        Ok(TestRow {
359            row_id,
360            value: decode_i64(bytes)?,
361        })
362    }
363
364    #[test]
365    fn spill_run_writes_and_reader_reads_entries_round_trip() {
366        let dir = tempfile::tempdir().unwrap();
367        let policy = MemoryPolicy::new(
368            Some(1),
369            SpillPolicy::SpillToDisk {
370                directory: dir.path().to_path_buf(),
371            },
372        );
373        let mut entries = vec![
374            (
375                TestRow {
376                    row_id: 2,
377                    value: 20,
378                },
379                20,
380            ),
381            (
382                TestRow {
383                    row_id: 1,
384                    value: 10,
385                },
386                10,
387            ),
388        ];
389
390        let path = spill_run(
391            &mut entries,
392            &policy,
393            "test-run",
394            |left, right| left.cmp(right),
395            |row| row.row_id,
396            encode_i64,
397            encode_row,
398        )
399        .unwrap();
400
401        assert!(entries.is_empty());
402        let mut reader = SpillRunReader::open(path.clone(), decode_i64, decode_row).unwrap();
403        assert_eq!(
404            reader
405                .next_entry()
406                .unwrap()
407                .map(|entry| (entry.row, entry.keys)),
408            Some((
409                TestRow {
410                    row_id: 1,
411                    value: 10
412                },
413                10
414            ))
415        );
416        assert_eq!(
417            reader
418                .next_entry()
419                .unwrap()
420                .map(|entry| (entry.row, entry.keys)),
421            Some((
422                TestRow {
423                    row_id: 2,
424                    value: 20
425                },
426                20
427            ))
428        );
429        assert!(reader.next_entry().unwrap().is_none());
430        drop(reader);
431        assert!(!path.exists());
432    }
433
434    #[test]
435    fn spill_merge_iterator_outputs_k_way_sorted_order() {
436        let dir = tempfile::tempdir().unwrap();
437        let policy = MemoryPolicy::new(
438            Some(1),
439            SpillPolicy::SpillToDisk {
440                directory: dir.path().to_path_buf(),
441            },
442        );
443        let mut left = vec![
444            (
445                TestRow {
446                    row_id: 3,
447                    value: 30,
448                },
449                30,
450            ),
451            (
452                TestRow {
453                    row_id: 1,
454                    value: 10,
455                },
456                10,
457            ),
458        ];
459        let mut right = vec![
460            (
461                TestRow {
462                    row_id: 4,
463                    value: 40,
464                },
465                40,
466            ),
467            (
468                TestRow {
469                    row_id: 2,
470                    value: 20,
471                },
472                20,
473            ),
474        ];
475        let run_a = spill_run(
476            &mut left,
477            &policy,
478            "test-run",
479            |left, right| left.cmp(right),
480            |row| row.row_id,
481            encode_i64,
482            encode_row,
483        )
484        .unwrap();
485        let run_b = spill_run(
486            &mut right,
487            &policy,
488            "test-run",
489            |left, right| left.cmp(right),
490            |row| row.row_id,
491            encode_i64,
492            encode_row,
493        )
494        .unwrap();
495
496        let mut iter = SpillMergeIterator::new(
497            vec![run_a, run_b],
498            |left, right| left.cmp(right),
499            |row: &TestRow| row.row_id,
500            decode_i64,
501            decode_row,
502        )
503        .unwrap();
504
505        let mut values = Vec::new();
506        while let Some(row) = iter.next_item() {
507            values.push(row.unwrap().value);
508        }
509
510        assert_eq!(values, vec![10, 20, 30, 40]);
511    }
512
513    #[test]
514    fn spill_failure_returns_stable_error_variant() {
515        let err = spill_io_error("sort spill", std::io::Error::other("disk full"));
516
517        assert!(matches!(err, Error::SpillFailed { .. }));
518        assert!(err.to_string().contains("sort spill"));
519        assert!(err.to_string().contains("disk full"));
520    }
521}