use crate::storage::{Page, PAGE_SIZE};
use crate::types::{PageId, Result, VelociError};
use async_trait::async_trait;
use parking_lot::RwLock;
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::fs::File;
use tokio::io::{AsyncReadExt, AsyncSeekExt, AsyncWriteExt};
use tokio::sync::RwLock as TokioRwLock;
#[async_trait]
pub trait AsyncVfs: Send + Sync {
async fn read_page(&self, page_id: PageId) -> Result<Page>;
async fn write_page(&self, page_id: PageId, page: &Page) -> Result<()>;
async fn allocate_page(&self) -> Result<PageId>;
async fn flush(&self) -> Result<()>;
async fn num_pages(&self) -> u64;
async fn sync(&self) -> Result<()>;
}
pub struct TokioVfs {
file_path: PathBuf,
file: Arc<TokioRwLock<Option<File>>>,
num_pages: Arc<RwLock<u64>>,
}
impl TokioVfs {
pub async fn new(path: impl AsRef<Path>) -> Result<Self> {
let path = path.as_ref().to_path_buf();
let file = tokio::fs::OpenOptions::new()
.read(true)
.write(true)
.create(true)
.open(&path)
.await?;
let metadata = file.metadata().await?;
let file_size = metadata.len();
let num_pages = if file_size == 0 {
0
} else {
(file_size + PAGE_SIZE as u64 - 1) / PAGE_SIZE as u64
};
Ok(Self {
file_path: path,
file: Arc::new(TokioRwLock::new(Some(file))),
num_pages: Arc::new(RwLock::new(num_pages)),
})
}
async fn get_file(&self) -> Result<File> {
let has_file = self.file.read().await.is_some();
if !has_file {
let file = tokio::fs::OpenOptions::new()
.read(true)
.write(true)
.open(&self.file_path)
.await?;
let file_clone = file.try_clone().await?;
*self.file.write().await = Some(file);
Ok(file_clone)
} else {
let guard = self.file.read().await;
guard.as_ref().unwrap().try_clone().await.map_err(|e| e.into())
}
}
}
#[async_trait]
impl AsyncVfs for TokioVfs {
async fn read_page(&self, page_id: PageId) -> Result<Page> {
let num_pages = *self.num_pages.read();
if page_id >= num_pages {
return Err(VelociError::NotFound(format!(
"Page {} out of bounds",
page_id
)));
}
let mut file = self.get_file().await?;
let mut page = Page::new();
let offset = page_id * PAGE_SIZE as u64;
file.seek(std::io::SeekFrom::Start(offset)).await?;
file.read_exact(page.data_mut()).await?;
Ok(page)
}
async fn write_page(&self, page_id: PageId, page: &Page) -> Result<()> {
let mut file = self.get_file().await?;
let offset = page_id * PAGE_SIZE as u64;
file.seek(std::io::SeekFrom::Start(offset)).await?;
file.write_all(page.data()).await?;
file.flush().await?;
let mut num_pages = self.num_pages.write();
if page_id >= *num_pages {
*num_pages = page_id + 1;
}
Ok(())
}
async fn allocate_page(&self) -> Result<PageId> {
let page_id = {
let mut num_pages = self.num_pages.write();
let page_id = *num_pages;
*num_pages += 1;
page_id
};
let page = Page::new();
self.write_page(page_id, &page).await?;
Ok(page_id)
}
async fn flush(&self) -> Result<()> {
let mut file = self.get_file().await?;
file.flush().await?;
Ok(())
}
async fn num_pages(&self) -> u64 {
*self.num_pages.read()
}
async fn sync(&self) -> Result<()> {
let file = self.get_file().await?;
file.sync_all().await?;
Ok(())
}
}
pub struct AsyncPageCache {
cache: Arc<RwLock<lru::LruCache<PageId, Arc<Page>>>>,
vfs: Arc<dyn AsyncVfs>,
}
impl AsyncPageCache {
pub fn new(capacity: usize, vfs: Arc<dyn AsyncVfs>) -> Self {
Self {
cache: Arc::new(RwLock::new(
lru::LruCache::new(std::num::NonZeroUsize::new(capacity).unwrap()),
)),
vfs,
}
}
pub async fn read_page(&self, page_id: PageId) -> Result<Arc<Page>> {
{
let mut cache = self.cache.write();
if let Some(page) = cache.get(&page_id) {
return Ok(Arc::clone(page));
}
}
let page = self.vfs.read_page(page_id).await?;
let page_arc = Arc::new(page);
let mut cache = self.cache.write();
cache.put(page_id, Arc::clone(&page_arc));
Ok(page_arc)
}
pub async fn write_page(&self, page_id: PageId, page: Page) -> Result<()> {
self.vfs.write_page(page_id, &page).await?;
let page_arc = Arc::new(page);
let mut cache = self.cache.write();
cache.put(page_id, page_arc);
Ok(())
}
pub async fn allocate_page(&self) -> Result<PageId> {
self.vfs.allocate_page().await
}
pub async fn flush(&self) -> Result<()> {
self.vfs.flush().await
}
pub fn cache_stats(&self) -> CacheStats {
let cache = self.cache.read();
CacheStats {
size: cache.len(),
capacity: cache.cap().get(),
}
}
}
#[derive(Debug, Clone)]
pub struct CacheStats {
pub size: usize,
pub capacity: usize,
}
pub struct AsyncPager {
cache: Arc<AsyncPageCache>,
vfs: Arc<dyn AsyncVfs>,
}
impl AsyncPager {
pub async fn new(path: impl AsRef<Path>, cache_size: usize) -> Result<Self> {
let vfs: Arc<dyn AsyncVfs> = Arc::new(TokioVfs::new(path).await?);
let cache = Arc::new(AsyncPageCache::new(cache_size, Arc::clone(&vfs)));
Ok(Self { cache, vfs })
}
pub async fn read_page(&self, page_id: PageId) -> Result<Arc<Page>> {
self.cache.read_page(page_id).await
}
pub async fn write_page(&self, page_id: PageId, page: Page) -> Result<()> {
self.cache.write_page(page_id, page).await
}
pub async fn allocate_page(&self) -> Result<PageId> {
self.cache.allocate_page().await
}
pub async fn flush(&self) -> Result<()> {
self.cache.flush().await?;
self.vfs.sync().await
}
pub async fn num_pages(&self) -> u64 {
self.vfs.num_pages().await
}
pub fn cache_stats(&self) -> CacheStats {
self.cache.cache_stats()
}
}
#[derive(Debug)]
pub enum IoRequest {
Read { page_id: PageId },
Write { page_id: PageId, page: Page },
}
pub struct BatchIoExecutor {
pager: Arc<AsyncPager>,
}
impl BatchIoExecutor {
pub fn new(pager: Arc<AsyncPager>) -> Self {
Self { pager }
}
pub async fn execute_batch(&self, requests: Vec<IoRequest>) -> Result<HashMap<PageId, Page>> {
let mut handles = Vec::new();
let mut results = HashMap::new();
for request in requests {
let pager = Arc::clone(&self.pager);
let handle = tokio::spawn(async move {
match request {
IoRequest::Read { page_id } => {
let page = pager.read_page(page_id).await?;
Ok::<_, VelociError>((page_id, (*page).clone()))
}
IoRequest::Write { page_id, page } => {
pager.write_page(page_id, page.clone()).await?;
Ok((page_id, page))
}
}
});
handles.push(handle);
}
for handle in handles {
match handle.await {
Ok(Ok((page_id, page))) => {
results.insert(page_id, page);
}
Ok(Err(e)) => return Err(e),
Err(e) => return Err(VelociError::IoError(e.to_string())),
}
}
Ok(results)
}
pub async fn read_batch(&self, page_ids: Vec<PageId>) -> Result<HashMap<PageId, Page>> {
let requests: Vec<IoRequest> = page_ids
.into_iter()
.map(|page_id| IoRequest::Read { page_id })
.collect();
self.execute_batch(requests).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
#[tokio::test]
async fn test_async_vfs_basic() {
let temp_file = NamedTempFile::new().unwrap();
let vfs = TokioVfs::new(temp_file.path()).await.unwrap();
let page_id = vfs.allocate_page().await.unwrap();
assert_eq!(page_id, 0);
let mut page = Page::new();
page.data_mut()[0..4].copy_from_slice(&[1, 2, 3, 4]);
vfs.write_page(page_id, &page).await.unwrap();
let read_page = vfs.read_page(page_id).await.unwrap();
assert_eq!(read_page.data()[0..4], [1, 2, 3, 4]);
}
#[tokio::test]
async fn test_async_page_cache() {
let temp_file = NamedTempFile::new().unwrap();
let vfs: Arc<dyn AsyncVfs> = Arc::new(TokioVfs::new(temp_file.path()).await.unwrap());
let cache = AsyncPageCache::new(10, vfs);
let page_id = cache.allocate_page().await.unwrap();
let mut page = Page::new();
page.data_mut()[0..4].copy_from_slice(&[5, 6, 7, 8]);
cache.write_page(page_id, page).await.unwrap();
let cached_page = cache.read_page(page_id).await.unwrap();
assert_eq!(cached_page.data()[0..4], [5, 6, 7, 8]);
let stats = cache.cache_stats();
assert_eq!(stats.size, 1);
assert_eq!(stats.capacity, 10);
}
#[tokio::test]
async fn test_batch_io() {
let temp_file = NamedTempFile::new().unwrap();
let pager = Arc::new(AsyncPager::new(temp_file.path(), 100).await.unwrap());
let executor = BatchIoExecutor::new(Arc::clone(&pager));
let page_ids: Vec<PageId> = futures::future::join_all(
(0..5).map(|_| pager.allocate_page())
)
.await
.into_iter()
.collect::<Result<Vec<_>>>()
.unwrap();
let mut write_requests = Vec::new();
for (i, &page_id) in page_ids.iter().enumerate() {
let mut page = Page::new();
page.data_mut()[0] = i as u8;
write_requests.push(IoRequest::Write {
page_id,
page,
});
}
executor.execute_batch(write_requests).await.unwrap();
let results = executor.read_batch(page_ids.clone()).await.unwrap();
assert_eq!(results.len(), 5);
for (i, &page_id) in page_ids.iter().enumerate() {
let page = results.get(&page_id).unwrap();
assert_eq!(page.data()[0], i as u8);
}
}
}