use std::io;
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use async_trait::async_trait;
use object_store::path::Path;
use tokio::io::AsyncWrite;
use lance_core::{Error, Result};
use crate::object_store::ObjectStore;
use crate::object_writer::WriteResult;
use crate::traits::{Reader, Writer};
#[async_trait]
pub trait SpillStore: Send + Sync + 'static {
async fn new_spill(&self) -> Result<(Box<dyn Writer>, Box<dyn Spill>)>;
}
#[async_trait]
pub trait Spill: Send + Sync {
async fn reader(&self) -> Result<Box<dyn Reader>>;
}
#[derive(Debug, Clone)]
struct DiskQuota {
cap_bytes: u64,
used: Arc<Mutex<u64>>,
}
impl DiskQuota {
fn new(cap_bytes: u64) -> Self {
Self {
cap_bytes,
used: Arc::new(Mutex::new(0)),
}
}
fn try_reserve(&self, n: u64) -> Result<()> {
let mut used = self.used.lock().unwrap();
let next = used.saturating_add(n);
if next > self.cap_bytes {
return Err(Error::disk_cap_exceeded(self.cap_bytes, *used));
}
*used = next;
Ok(())
}
fn release(&self, n: u64) {
let mut used = self.used.lock().unwrap();
*used = used.saturating_sub(n);
}
}
struct SpillWriter {
inner: Box<dyn Writer>,
quota: Option<DiskQuota>,
finished: Arc<AtomicBool>,
}
impl AsyncWrite for SpillWriter {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let this = self.get_mut();
let Some(quota) = &this.quota else {
return Pin::new(this.inner.as_mut()).poll_write(cx, buf);
};
if let Err(e) = quota.try_reserve(buf.len() as u64) {
return Poll::Ready(Err(io::Error::other(e)));
}
let poll = Pin::new(this.inner.as_mut()).poll_write(cx, buf);
match &poll {
Poll::Ready(Ok(n)) => quota.release((buf.len() - *n) as u64),
_ => quota.release(buf.len() as u64),
}
poll
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(self.get_mut().inner.as_mut()).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
let poll = Pin::new(this.inner.as_mut()).poll_shutdown(cx);
if matches!(poll, Poll::Ready(Ok(()))) {
this.finished.store(true, Ordering::Relaxed);
}
poll
}
}
#[async_trait]
impl Writer for SpillWriter {
async fn tell(&mut self) -> Result<usize> {
self.inner.tell().await
}
async fn shutdown(&mut self) -> Result<WriteResult> {
let result = self.inner.shutdown().await?;
self.finished.store(true, Ordering::Relaxed);
Ok(result)
}
}
pub struct LocalSpillStore {
store: Arc<ObjectStore>,
temp_dir: Arc<tempfile::TempDir>,
file_counter: Arc<AtomicU64>,
quota: Option<DiskQuota>,
}
impl LocalSpillStore {
pub fn new() -> Result<Self> {
Ok(Self {
store: Arc::new(ObjectStore::local()),
temp_dir: Arc::new(tempfile::tempdir()?),
file_counter: Arc::new(AtomicU64::new(0)),
quota: None,
})
}
pub fn with_cap(cap_bytes: u64) -> Result<Self> {
Ok(Self {
store: Arc::new(ObjectStore::local()),
temp_dir: Arc::new(tempfile::tempdir()?),
file_counter: Arc::new(AtomicU64::new(0)),
quota: Some(DiskQuota::new(cap_bytes)),
})
}
}
impl Default for LocalSpillStore {
fn default() -> Self {
Self::new().expect("failed to create temp directory for LocalSpillStore")
}
}
#[async_trait]
impl SpillStore for LocalSpillStore {
async fn new_spill(&self) -> Result<(Box<dyn Writer>, Box<dyn Spill>)> {
let idx = self.file_counter.fetch_add(1, Ordering::Relaxed);
let fs_path = self.temp_dir.path().join(format!("spill_{idx:06}.bin"));
let os_path = Path::from_absolute_path(&fs_path)?;
let finished = Arc::new(AtomicBool::new(false));
let writer = Box::new(SpillWriter {
inner: self.store.create(&os_path).await?,
quota: self.quota.clone(),
finished: finished.clone(),
});
let spill = Box::new(LocalSpill {
store: self.store.clone(),
os_path,
fs_path,
quota: self.quota.clone(),
finished,
_temp_dir: self.temp_dir.clone(),
});
Ok((writer, spill))
}
}
struct LocalSpill {
store: Arc<ObjectStore>,
os_path: Path,
fs_path: PathBuf,
quota: Option<DiskQuota>,
finished: Arc<AtomicBool>,
_temp_dir: Arc<tempfile::TempDir>,
}
#[async_trait]
impl Spill for LocalSpill {
async fn reader(&self) -> Result<Box<dyn Reader>> {
if !self.finished.load(Ordering::Relaxed) {
return Err(Error::invalid_input(
"spill reader requested before the writer was shut down",
));
}
self.store.open(&self.os_path).await
}
}
impl Drop for LocalSpill {
fn drop(&mut self) {
if let Some(quota) = &self.quota
&& let Ok(metadata) = std::fs::metadata(&self.fs_path)
{
quota.release(metadata.len());
}
let _ = std::fs::remove_file(&self.fs_path);
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::AsyncWriteExt;
async fn finish_writer(mut writer: Box<dyn Writer>, data: &[u8]) -> Result<()> {
writer.write_all(data).await?;
Writer::shutdown(writer.as_mut()).await?;
Ok(())
}
#[test]
fn test_disk_quota_reserve_release() {
let quota = DiskQuota::new(100);
quota.try_reserve(60).unwrap();
assert!(quota.try_reserve(60).is_err());
quota.release(60);
quota.try_reserve(60).unwrap();
quota.try_reserve(40).unwrap();
assert!(quota.try_reserve(1).is_err());
}
#[tokio::test]
async fn test_write_then_read() {
let store = LocalSpillStore::new().unwrap();
let (writer, spill) = store.new_spill().await.unwrap();
let data = b"hello spill world";
finish_writer(writer, data).await.unwrap();
let reader = spill.reader().await.unwrap();
let read_back = reader.get_all().await.unwrap();
assert_eq!(read_back.as_ref(), data);
}
#[tokio::test]
async fn test_reader_requires_finished_writer() {
let store = LocalSpillStore::new().unwrap();
let (mut writer, spill) = store.new_spill().await.unwrap();
writer.write_all(b"partial").await.unwrap();
let Err(err) = spill.reader().await else {
panic!("reader before shutdown should be rejected");
};
assert!(
matches!(err, Error::InvalidInput { .. }),
"expected InvalidInput, got {err:?}"
);
Writer::shutdown(writer.as_mut()).await.unwrap();
let reader = spill.reader().await.unwrap();
assert_eq!(reader.get_all().await.unwrap().as_ref(), b"partial");
}
#[tokio::test]
async fn test_reader_ready_after_async_shutdown() {
let store = LocalSpillStore::new().unwrap();
let (mut writer, spill) = store.new_spill().await.unwrap();
writer.write_all(b"async").await.unwrap();
AsyncWriteExt::shutdown(&mut writer).await.unwrap();
let reader = spill.reader().await.unwrap();
assert_eq!(reader.get_all().await.unwrap().as_ref(), b"async");
}
#[tokio::test]
async fn test_empty_spill() {
let store = LocalSpillStore::with_cap(100).unwrap();
let (writer, spill) = store.new_spill().await.unwrap();
finish_writer(writer, b"").await.unwrap();
let reader = spill.reader().await.unwrap();
assert!(reader.get_all().await.unwrap().is_empty());
}
#[tokio::test]
async fn test_raii_cleanup() {
let store = LocalSpillStore::new().unwrap();
let (writer, spill) = store.new_spill().await.unwrap();
finish_writer(writer, b"some bytes").await.unwrap();
let path = store.temp_dir.path().join("spill_000000.bin");
assert!(path.exists());
drop(spill);
assert!(!path.exists(), "spill file should be deleted on drop");
}
#[tokio::test]
async fn test_cap_exceeded() {
let store = LocalSpillStore::with_cap(100).unwrap();
let (writer, _spill) = store.new_spill().await.unwrap();
let err = finish_writer(writer, &[0u8; 101]).await.unwrap_err();
assert!(
matches!(err, Error::DiskCapExceeded { cap_bytes: 100, .. }),
"expected DiskCapExceeded, got {err:?}"
);
}
#[tokio::test]
async fn test_cap_shared_across_files() {
let store = LocalSpillStore::with_cap(100).unwrap();
let (writer_a, _spill_a) = store.new_spill().await.unwrap();
let (writer_b, _spill_b) = store.new_spill().await.unwrap();
finish_writer(writer_a, &[0u8; 60]).await.unwrap();
let err = finish_writer(writer_b, &[0u8; 60]).await.unwrap_err();
assert!(
matches!(err, Error::DiskCapExceeded { cap_bytes: 100, .. }),
"expected DiskCapExceeded, got {err:?}"
);
}
#[tokio::test]
async fn test_cap_freed_on_drop() {
let store = LocalSpillStore::with_cap(100).unwrap();
{
let (writer, spill) = store.new_spill().await.unwrap();
finish_writer(writer, &[0u8; 80]).await.unwrap();
drop(spill);
}
let (writer, _spill) = store.new_spill().await.unwrap();
finish_writer(writer, &[0u8; 80]).await.unwrap();
}
#[tokio::test]
async fn test_custom_implementation() {
struct MemStore;
struct MemSpill;
#[async_trait]
impl Spill for MemSpill {
async fn reader(&self) -> Result<Box<dyn Reader>> {
ObjectStore::memory().open(&Path::from("/mem")).await
}
}
#[async_trait]
impl SpillStore for MemStore {
async fn new_spill(&self) -> Result<(Box<dyn Writer>, Box<dyn Spill>)> {
let writer = ObjectStore::memory().create(&Path::from("/mem")).await?;
Ok((writer, Box::new(MemSpill)))
}
}
let store = MemStore;
let (_writer, _spill) = store.new_spill().await.unwrap();
}
struct ControlledWriter {
outcome: Poll<io::Result<usize>>,
}
impl AsyncWrite for ControlledWriter {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
match &self.outcome {
Poll::Ready(Ok(n)) => Poll::Ready(Ok((*n).min(buf.len()))),
Poll::Ready(Err(e)) => Poll::Ready(Err(io::Error::new(e.kind(), e.to_string()))),
Poll::Pending => Poll::Pending,
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
#[async_trait]
impl Writer for ControlledWriter {
async fn tell(&mut self) -> Result<usize> {
Ok(0)
}
async fn shutdown(&mut self) -> Result<WriteResult> {
Ok(WriteResult::default())
}
}
#[tokio::test]
async fn test_spill_writer_releases_unaccepted_bytes() {
let quota = DiskQuota::new(100);
let mut writer = SpillWriter {
inner: Box::new(ControlledWriter {
outcome: Poll::Ready(Ok(10)),
}),
quota: Some(quota.clone()),
finished: Arc::new(AtomicBool::new(false)),
};
let n = writer.write(&[0u8; 40]).await.unwrap();
assert_eq!(n, 10);
assert_eq!(
*quota.used.lock().unwrap(),
10,
"only the accepted bytes should remain reserved"
);
let quota = DiskQuota::new(100);
let mut writer = SpillWriter {
inner: Box::new(ControlledWriter {
outcome: Poll::Ready(Err(io::Error::other("boom"))),
}),
quota: Some(quota.clone()),
finished: Arc::new(AtomicBool::new(false)),
};
writer.write(&[0u8; 40]).await.unwrap_err();
assert_eq!(
*quota.used.lock().unwrap(),
0,
"a failed write should release its entire reservation"
);
}
}