use super::pager::{Page, PageId, Pager};
use lru::LruCache;
use std::num::NonZeroUsize;
pub struct BufferPool {
pager: Pager,
cache: LruCache<PageId, (Page, bool)>,
dirty_write_count: u32,
sync_interval: u32,
}
impl BufferPool {
pub fn new(pager: Pager, capacity: usize) -> Self {
let cap = NonZeroUsize::new(capacity.max(1)).unwrap_or(NonZeroUsize::MIN);
Self {
pager,
cache: LruCache::new(cap),
dirty_write_count: 0,
sync_interval: 0,
}
}
pub fn with_sync_interval(mut self, interval: u32) -> Self {
self.sync_interval = interval;
self
}
pub fn fetch_page(&mut self, page_id: PageId) -> anyhow::Result<Page> {
if let Some((page, _)) = self.cache.get(&page_id) {
return Ok(page.clone());
}
let page = self.pager.read_page(page_id)?;
self.put_and_evict(page_id, page.clone(), false)?;
Ok(page)
}
pub fn write_page(&mut self, page_id: PageId, page: &Page) -> anyhow::Result<()> {
self.put_and_evict(page_id, page.clone(), true)?;
if self.sync_interval > 0 {
self.dirty_write_count += 1;
if self.dirty_write_count >= self.sync_interval {
self.flush_and_sync()?;
self.dirty_write_count = 0;
}
}
Ok(())
}
pub fn allocate_page(&mut self) -> anyhow::Result<PageId> {
let page_id = self.pager.allocate_page()?;
Ok(page_id)
}
pub fn free_page(&mut self, page_id: PageId) -> anyhow::Result<()> {
self.cache.pop(&page_id);
self.pager.free_page(page_id)?;
Ok(())
}
fn put_and_evict(&mut self, page_id: PageId, page: Page, is_dirty: bool) -> anyhow::Result<()> {
if self.cache.len() == self.cache.cap().get() {
if !self.cache.contains(&page_id) {
if let Some((evict_id, (evict_page, evict_dirty))) = self.cache.pop_lru() {
if evict_dirty {
self.pager.write_page(evict_id, &evict_page)?;
}
}
}
}
if let Some((_, curr_dirty)) = self.cache.get(&page_id) {
let new_dirty = is_dirty || *curr_dirty;
self.cache.put(page_id, (page, new_dirty));
} else {
self.cache.put(page_id, (page, is_dirty));
}
Ok(())
}
pub fn flush_all(&mut self) -> anyhow::Result<()> {
let mut to_flush = Vec::new();
for (page_id, (page, dirty)) in self.cache.iter() {
if *dirty {
to_flush.push((*page_id, page.clone()));
}
}
for (page_id, page) in to_flush {
self.pager.write_page(page_id, &page)?;
if let Some(entry) = self.cache.get_mut(&page_id) {
entry.1 = false;
}
}
Ok(())
}
pub fn flush_and_sync(&mut self) -> anyhow::Result<()> {
self.flush_all()?;
self.pager.sync()?;
Ok(())
}
pub fn sync(&mut self) -> anyhow::Result<()> {
self.pager.sync()?;
Ok(())
}
pub fn get_num_pages(&self) -> u32 {
self.pager.num_pages
}
pub fn pager_mut(&mut self) -> &mut Pager {
&mut self.pager
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
#[test]
fn test_buffer_pool_eviction() -> anyhow::Result<()> {
let temp_file = NamedTempFile::new()?;
let path = temp_file.path().to_path_buf();
let pager = Pager::open(&path)?;
let mut pool = BufferPool::new(pager, 2);
let id0 = pool.allocate_page()?;
let id1 = pool.allocate_page()?;
let id2 = pool.allocate_page()?;
let mut page0 = Page::default();
page0.data[0] = 100;
pool.write_page(id0, &page0)?;
let mut page1 = Page::default();
page1.data[0] = 101;
pool.write_page(id1, &page1)?;
let mut page2 = Page::default();
page2.data[0] = 102;
pool.write_page(id2, &page2)?;
let mut direct_pager = Pager::open(&path)?;
let read_page0 = direct_pager.read_page(id0)?;
assert_eq!(read_page0.data[0], 100);
Ok(())
}
#[test]
fn test_buffer_pool_flush_and_sync() -> anyhow::Result<()> {
let temp_file = NamedTempFile::new()?;
let path = temp_file.path().to_path_buf();
let pager = Pager::open(&path)?;
let mut pool = BufferPool::new(pager, 10);
let id = pool.allocate_page()?;
let mut page = Page::default();
page.data[0] = 77;
pool.write_page(id, &page)?;
pool.flush_and_sync()?;
let mut direct_pager = Pager::open(&path)?;
let read_page = direct_pager.read_page(id)?;
assert_eq!(read_page.data[0], 77);
Ok(())
}
}