1use std::io;
39use std::path::PathBuf;
40use std::pin::Pin;
41use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
42use std::sync::{Arc, Mutex};
43use std::task::{Context, Poll};
44
45use async_trait::async_trait;
46use object_store::path::Path;
47use tokio::io::AsyncWrite;
48
49use lance_core::{Error, Result};
50
51use crate::object_store::ObjectStore;
52use crate::object_writer::WriteResult;
53use crate::traits::{Reader, Writer};
54
55#[async_trait]
61pub trait SpillStore: Send + Sync + 'static {
62 async fn new_spill(&self) -> Result<(Box<dyn Writer>, Box<dyn Spill>)>;
69}
70
71#[async_trait]
76pub trait Spill: Send + Sync {
77 async fn reader(&self) -> Result<Box<dyn Reader>>;
82}
83
84#[derive(Debug, Clone)]
89struct DiskQuota {
90 cap_bytes: u64,
91 used: Arc<Mutex<u64>>,
92}
93
94impl DiskQuota {
95 fn new(cap_bytes: u64) -> Self {
96 Self {
97 cap_bytes,
98 used: Arc::new(Mutex::new(0)),
99 }
100 }
101
102 fn try_reserve(&self, n: u64) -> Result<()> {
105 let mut used = self.used.lock().unwrap();
108 let next = used.saturating_add(n);
109 if next > self.cap_bytes {
110 return Err(Error::disk_cap_exceeded(self.cap_bytes, *used));
111 }
112 *used = next;
113 Ok(())
114 }
115
116 fn release(&self, n: u64) {
118 let mut used = self.used.lock().unwrap();
120 *used = used.saturating_sub(n);
121 }
122}
123
124struct SpillWriter {
131 inner: Box<dyn Writer>,
132 quota: Option<DiskQuota>,
133 finished: Arc<AtomicBool>,
134}
135
136impl AsyncWrite for SpillWriter {
137 fn poll_write(
138 self: Pin<&mut Self>,
139 cx: &mut Context<'_>,
140 buf: &[u8],
141 ) -> Poll<io::Result<usize>> {
142 let this = self.get_mut();
143 let Some(quota) = &this.quota else {
144 return Pin::new(this.inner.as_mut()).poll_write(cx, buf);
145 };
146 if let Err(e) = quota.try_reserve(buf.len() as u64) {
150 return Poll::Ready(Err(io::Error::other(e)));
151 }
152 let poll = Pin::new(this.inner.as_mut()).poll_write(cx, buf);
153 match &poll {
154 Poll::Ready(Ok(n)) => quota.release((buf.len() - *n) as u64),
155 _ => quota.release(buf.len() as u64),
156 }
157 poll
158 }
159
160 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
161 Pin::new(self.get_mut().inner.as_mut()).poll_flush(cx)
162 }
163
164 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
165 let this = self.get_mut();
166 let poll = Pin::new(this.inner.as_mut()).poll_shutdown(cx);
167 if matches!(poll, Poll::Ready(Ok(()))) {
168 this.finished.store(true, Ordering::Relaxed);
171 }
172 poll
173 }
174}
175
176#[async_trait]
177impl Writer for SpillWriter {
178 async fn tell(&mut self) -> Result<usize> {
179 self.inner.tell().await
180 }
181
182 async fn shutdown(&mut self) -> Result<WriteResult> {
183 let result = self.inner.shutdown().await?;
184 self.finished.store(true, Ordering::Relaxed);
188 Ok(result)
189 }
190}
191
192pub struct LocalSpillStore {
200 store: Arc<ObjectStore>,
201 temp_dir: Arc<tempfile::TempDir>,
203 file_counter: Arc<AtomicU64>,
204 quota: Option<DiskQuota>,
206}
207
208impl LocalSpillStore {
209 pub fn new() -> Result<Self> {
211 Ok(Self {
212 store: Arc::new(ObjectStore::local()),
213 temp_dir: Arc::new(tempfile::tempdir()?),
214 file_counter: Arc::new(AtomicU64::new(0)),
215 quota: None,
216 })
217 }
218
219 pub fn with_cap(cap_bytes: u64) -> Result<Self> {
222 Ok(Self {
223 store: Arc::new(ObjectStore::local()),
224 temp_dir: Arc::new(tempfile::tempdir()?),
225 file_counter: Arc::new(AtomicU64::new(0)),
226 quota: Some(DiskQuota::new(cap_bytes)),
227 })
228 }
229}
230
231impl Default for LocalSpillStore {
232 fn default() -> Self {
233 Self::new().expect("failed to create temp directory for LocalSpillStore")
234 }
235}
236
237#[async_trait]
238impl SpillStore for LocalSpillStore {
239 async fn new_spill(&self) -> Result<(Box<dyn Writer>, Box<dyn Spill>)> {
240 let idx = self.file_counter.fetch_add(1, Ordering::Relaxed);
241 let fs_path = self.temp_dir.path().join(format!("spill_{idx:06}.bin"));
242 let os_path = Path::from_absolute_path(&fs_path)?;
243 let finished = Arc::new(AtomicBool::new(false));
244
245 let writer = Box::new(SpillWriter {
246 inner: self.store.create(&os_path).await?,
247 quota: self.quota.clone(),
248 finished: finished.clone(),
249 });
250 let spill = Box::new(LocalSpill {
251 store: self.store.clone(),
252 os_path,
253 fs_path,
254 quota: self.quota.clone(),
255 finished,
256 _temp_dir: self.temp_dir.clone(),
257 });
258 Ok((writer, spill))
259 }
260}
261
262struct LocalSpill {
264 store: Arc<ObjectStore>,
265 os_path: Path,
266 fs_path: PathBuf,
267 quota: Option<DiskQuota>,
268 finished: Arc<AtomicBool>,
270 _temp_dir: Arc<tempfile::TempDir>,
272}
273
274#[async_trait]
275impl Spill for LocalSpill {
276 async fn reader(&self) -> Result<Box<dyn Reader>> {
277 if !self.finished.load(Ordering::Relaxed) {
281 return Err(Error::invalid_input(
282 "spill reader requested before the writer was shut down",
283 ));
284 }
285 self.store.open(&self.os_path).await
286 }
287}
288
289impl Drop for LocalSpill {
290 fn drop(&mut self) {
291 if let Some(quota) = &self.quota
295 && let Ok(metadata) = std::fs::metadata(&self.fs_path)
296 {
297 quota.release(metadata.len());
298 }
299 let _ = std::fs::remove_file(&self.fs_path);
301 }
302}
303
304#[cfg(test)]
305mod tests {
306 use super::*;
307 use tokio::io::AsyncWriteExt;
308
309 async fn finish_writer(mut writer: Box<dyn Writer>, data: &[u8]) -> Result<()> {
311 writer.write_all(data).await?;
312 Writer::shutdown(writer.as_mut()).await?;
313 Ok(())
314 }
315
316 #[test]
317 fn test_disk_quota_reserve_release() {
318 let quota = DiskQuota::new(100);
319 quota.try_reserve(60).unwrap();
320 assert!(quota.try_reserve(60).is_err());
321 quota.release(60);
322 quota.try_reserve(60).unwrap();
323 quota.try_reserve(40).unwrap();
325 assert!(quota.try_reserve(1).is_err());
326 }
327
328 #[tokio::test]
329 async fn test_write_then_read() {
330 let store = LocalSpillStore::new().unwrap();
331 let (writer, spill) = store.new_spill().await.unwrap();
332
333 let data = b"hello spill world";
334 finish_writer(writer, data).await.unwrap();
335
336 let reader = spill.reader().await.unwrap();
337 let read_back = reader.get_all().await.unwrap();
338 assert_eq!(read_back.as_ref(), data);
339 }
340
341 #[tokio::test]
342 async fn test_reader_requires_finished_writer() {
343 let store = LocalSpillStore::new().unwrap();
344 let (mut writer, spill) = store.new_spill().await.unwrap();
345 writer.write_all(b"partial").await.unwrap();
346
347 let Err(err) = spill.reader().await else {
349 panic!("reader before shutdown should be rejected");
350 };
351 assert!(
352 matches!(err, Error::InvalidInput { .. }),
353 "expected InvalidInput, got {err:?}"
354 );
355
356 Writer::shutdown(writer.as_mut()).await.unwrap();
358 let reader = spill.reader().await.unwrap();
359 assert_eq!(reader.get_all().await.unwrap().as_ref(), b"partial");
360 }
361
362 #[tokio::test]
363 async fn test_reader_ready_after_async_shutdown() {
364 let store = LocalSpillStore::new().unwrap();
368 let (mut writer, spill) = store.new_spill().await.unwrap();
369 writer.write_all(b"async").await.unwrap();
370 AsyncWriteExt::shutdown(&mut writer).await.unwrap();
371
372 let reader = spill.reader().await.unwrap();
373 assert_eq!(reader.get_all().await.unwrap().as_ref(), b"async");
374 }
375
376 #[tokio::test]
377 async fn test_empty_spill() {
378 let store = LocalSpillStore::with_cap(100).unwrap();
381 let (writer, spill) = store.new_spill().await.unwrap();
382 finish_writer(writer, b"").await.unwrap();
383
384 let reader = spill.reader().await.unwrap();
385 assert!(reader.get_all().await.unwrap().is_empty());
386 }
387
388 #[tokio::test]
389 async fn test_raii_cleanup() {
390 let store = LocalSpillStore::new().unwrap();
391 let (writer, spill) = store.new_spill().await.unwrap();
392 finish_writer(writer, b"some bytes").await.unwrap();
393
394 let path = store.temp_dir.path().join("spill_000000.bin");
396 assert!(path.exists());
397 drop(spill);
398 assert!(!path.exists(), "spill file should be deleted on drop");
399 }
400
401 #[tokio::test]
402 async fn test_cap_exceeded() {
403 let store = LocalSpillStore::with_cap(100).unwrap();
404 let (writer, _spill) = store.new_spill().await.unwrap();
405 let err = finish_writer(writer, &[0u8; 101]).await.unwrap_err();
406 assert!(
407 matches!(err, Error::DiskCapExceeded { cap_bytes: 100, .. }),
408 "expected DiskCapExceeded, got {err:?}"
409 );
410 }
411
412 #[tokio::test]
413 async fn test_cap_shared_across_files() {
414 let store = LocalSpillStore::with_cap(100).unwrap();
415 let (writer_a, _spill_a) = store.new_spill().await.unwrap();
416 let (writer_b, _spill_b) = store.new_spill().await.unwrap();
417
418 finish_writer(writer_a, &[0u8; 60]).await.unwrap();
419 let err = finish_writer(writer_b, &[0u8; 60]).await.unwrap_err();
421 assert!(
422 matches!(err, Error::DiskCapExceeded { cap_bytes: 100, .. }),
423 "expected DiskCapExceeded, got {err:?}"
424 );
425 }
426
427 #[tokio::test]
428 async fn test_cap_freed_on_drop() {
429 let store = LocalSpillStore::with_cap(100).unwrap();
430
431 {
432 let (writer, spill) = store.new_spill().await.unwrap();
433 finish_writer(writer, &[0u8; 80]).await.unwrap();
434 drop(spill);
436 }
437
438 let (writer, _spill) = store.new_spill().await.unwrap();
439 finish_writer(writer, &[0u8; 80]).await.unwrap();
441 }
442
443 #[tokio::test]
444 async fn test_custom_implementation() {
445 struct MemStore;
447 struct MemSpill;
448
449 #[async_trait]
450 impl Spill for MemSpill {
451 async fn reader(&self) -> Result<Box<dyn Reader>> {
452 ObjectStore::memory().open(&Path::from("/mem")).await
453 }
454 }
455
456 #[async_trait]
457 impl SpillStore for MemStore {
458 async fn new_spill(&self) -> Result<(Box<dyn Writer>, Box<dyn Spill>)> {
459 let writer = ObjectStore::memory().create(&Path::from("/mem")).await?;
460 Ok((writer, Box::new(MemSpill)))
461 }
462 }
463
464 let store = MemStore;
465 let (_writer, _spill) = store.new_spill().await.unwrap();
468 }
469
470 struct ControlledWriter {
474 outcome: Poll<io::Result<usize>>,
475 }
476
477 impl AsyncWrite for ControlledWriter {
478 fn poll_write(
479 self: Pin<&mut Self>,
480 _cx: &mut Context<'_>,
481 buf: &[u8],
482 ) -> Poll<io::Result<usize>> {
483 match &self.outcome {
484 Poll::Ready(Ok(n)) => Poll::Ready(Ok((*n).min(buf.len()))),
485 Poll::Ready(Err(e)) => Poll::Ready(Err(io::Error::new(e.kind(), e.to_string()))),
486 Poll::Pending => Poll::Pending,
487 }
488 }
489 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
490 Poll::Ready(Ok(()))
491 }
492 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
493 Poll::Ready(Ok(()))
494 }
495 }
496
497 #[async_trait]
498 impl Writer for ControlledWriter {
499 async fn tell(&mut self) -> Result<usize> {
500 Ok(0)
501 }
502 async fn shutdown(&mut self) -> Result<WriteResult> {
503 Ok(WriteResult::default())
504 }
505 }
506
507 #[tokio::test]
508 async fn test_spill_writer_releases_unaccepted_bytes() {
509 let quota = DiskQuota::new(100);
512 let mut writer = SpillWriter {
513 inner: Box::new(ControlledWriter {
514 outcome: Poll::Ready(Ok(10)),
515 }),
516 quota: Some(quota.clone()),
517 finished: Arc::new(AtomicBool::new(false)),
518 };
519 let n = writer.write(&[0u8; 40]).await.unwrap();
520 assert_eq!(n, 10);
521 assert_eq!(
522 *quota.used.lock().unwrap(),
523 10,
524 "only the accepted bytes should remain reserved"
525 );
526
527 let quota = DiskQuota::new(100);
529 let mut writer = SpillWriter {
530 inner: Box::new(ControlledWriter {
531 outcome: Poll::Ready(Err(io::Error::other("boom"))),
532 }),
533 quota: Some(quota.clone()),
534 finished: Arc::new(AtomicBool::new(false)),
535 };
536 writer.write(&[0u8; 40]).await.unwrap_err();
537 assert_eq!(
538 *quota.used.lock().unwrap(),
539 0,
540 "a failed write should release its entire reservation"
541 );
542 }
543}