1use 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
21pub 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
89pub 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
95pub 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
115pub fn spill_io_error(operation: &str, err: impl std::fmt::Display) -> Error {
117 Error::SpillFailed {
118 reason: format!("{operation}: {err}"),
119 }
120}
121
122pub struct SpillEntry<T, K> {
124 pub row: T,
126 pub keys: K,
128}
129
130pub 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 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 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
204pub 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 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 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}