use crate::error::{Result, TdbError};
use crate::storage::file_manager::FileManager;
use crate::storage::page::{Page, PageId, PageType};
use dashmap::DashMap;
use parking_lot::RwLock;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
pub type FrameId = usize;
struct BufferFrame {
page: RwLock<Option<Page>>,
access_count: AtomicU64,
pin_count: AtomicU64,
}
impl BufferFrame {
fn new() -> Self {
Self {
page: RwLock::new(None),
access_count: AtomicU64::new(0),
pin_count: AtomicU64::new(0),
}
}
fn is_empty(&self) -> bool {
self.page.read().is_none()
}
fn is_pinned(&self) -> bool {
self.pin_count.load(Ordering::Acquire) > 0
}
fn pin(&self) {
self.pin_count.fetch_add(1, Ordering::AcqRel);
}
fn unpin(&self) {
let prev = self.pin_count.fetch_sub(1, Ordering::AcqRel);
debug_assert!(prev > 0, "Unpin called on unpinned frame");
}
fn access(&self) {
self.access_count.fetch_add(1, Ordering::AcqRel);
}
fn get_access_count(&self) -> u64 {
self.access_count.load(Ordering::Acquire)
}
}
pub struct BufferPool {
frames: Vec<BufferFrame>,
page_table: DashMap<PageId, FrameId>,
file_manager: Arc<FileManager>,
clock_hand: AtomicU64,
stats: BufferPoolStats,
}
#[derive(Debug, Default)]
pub struct BufferPoolStats {
pub total_fetches: AtomicU64,
pub cache_hits: AtomicU64,
pub cache_misses: AtomicU64,
pub evictions: AtomicU64,
pub writes: AtomicU64,
}
impl Clone for BufferPoolStats {
fn clone(&self) -> Self {
Self {
total_fetches: AtomicU64::new(self.total_fetches.load(Ordering::Relaxed)),
cache_hits: AtomicU64::new(self.cache_hits.load(Ordering::Relaxed)),
cache_misses: AtomicU64::new(self.cache_misses.load(Ordering::Relaxed)),
evictions: AtomicU64::new(self.evictions.load(Ordering::Relaxed)),
writes: AtomicU64::new(self.writes.load(Ordering::Relaxed)),
}
}
}
impl BufferPoolStats {
pub fn hit_rate(&self) -> f64 {
let total = self.total_fetches.load(Ordering::Relaxed);
if total == 0 {
0.0
} else {
let hits = self.cache_hits.load(Ordering::Relaxed);
hits as f64 / total as f64
}
}
pub fn reset(&self) {
self.total_fetches.store(0, Ordering::Relaxed);
self.cache_hits.store(0, Ordering::Relaxed);
self.cache_misses.store(0, Ordering::Relaxed);
self.evictions.store(0, Ordering::Relaxed);
self.writes.store(0, Ordering::Relaxed);
}
pub fn snapshot(&self) -> Self {
Self {
total_fetches: AtomicU64::new(self.total_fetches.load(Ordering::Relaxed)),
cache_hits: AtomicU64::new(self.cache_hits.load(Ordering::Relaxed)),
cache_misses: AtomicU64::new(self.cache_misses.load(Ordering::Relaxed)),
evictions: AtomicU64::new(self.evictions.load(Ordering::Relaxed)),
writes: AtomicU64::new(self.writes.load(Ordering::Relaxed)),
}
}
}
impl BufferPool {
pub fn new(pool_size: usize, file_manager: Arc<FileManager>) -> Self {
let mut frames = Vec::with_capacity(pool_size);
for _ in 0..pool_size {
frames.push(BufferFrame::new());
}
Self {
frames,
page_table: DashMap::new(),
file_manager,
clock_hand: AtomicU64::new(0),
stats: BufferPoolStats::default(),
}
}
pub fn fetch_page(&self, page_id: PageId) -> Result<PageGuard<'_>> {
self.stats.total_fetches.fetch_add(1, Ordering::Relaxed);
if let Some(frame_entry) = self.page_table.get(&page_id) {
let frame_id = *frame_entry;
let frame = &self.frames[frame_id];
frame.pin();
frame.access();
self.stats.cache_hits.fetch_add(1, Ordering::Relaxed);
return Ok(PageGuard {
frame_id,
page_id,
buffer_pool: self,
});
}
self.stats.cache_misses.fetch_add(1, Ordering::Relaxed);
let frame_id = self.find_victim_frame()?;
let frame = &self.frames[frame_id];
{
let mut page_guard = frame.page.write();
if let Some(old_page) = page_guard.take() {
if old_page.is_dirty() {
let old_page_id = old_page.page_id();
self.page_table.remove(&old_page_id);
drop(page_guard);
let mut write_page = old_page;
self.file_manager.write_page(&mut write_page)?;
self.stats.writes.fetch_add(1, Ordering::Relaxed);
page_guard = frame.page.write();
} else {
let old_page_id = old_page.page_id();
self.page_table.remove(&old_page_id);
}
self.stats.evictions.fetch_add(1, Ordering::Relaxed);
}
let page = self.file_manager.read_page(page_id)?;
*page_guard = Some(page);
}
self.page_table.insert(page_id, frame_id);
frame.pin();
frame.access();
Ok(PageGuard {
frame_id,
page_id,
buffer_pool: self,
})
}
pub fn new_page(&self, page_type: PageType) -> Result<PageGuard<'_>> {
let page_id = self.file_manager.allocate_page()?;
let frame_id = self.find_victim_frame()?;
let frame = &self.frames[frame_id];
{
let mut page_guard = frame.page.write();
if let Some(old_page) = page_guard.take() {
if old_page.is_dirty() {
let old_page_id = old_page.page_id();
self.page_table.remove(&old_page_id);
drop(page_guard);
let mut write_page = old_page;
self.file_manager.write_page(&mut write_page)?;
self.stats.writes.fetch_add(1, Ordering::Relaxed);
page_guard = frame.page.write();
} else {
let old_page_id = old_page.page_id();
self.page_table.remove(&old_page_id);
}
self.stats.evictions.fetch_add(1, Ordering::Relaxed);
}
let page = Page::new(page_id, page_type);
*page_guard = Some(page);
}
self.page_table.insert(page_id, frame_id);
frame.pin();
frame.access();
Ok(PageGuard {
frame_id,
page_id,
buffer_pool: self,
})
}
pub fn flush_page(&self, page_id: PageId) -> Result<()> {
if let Some(frame_entry) = self.page_table.get(&page_id) {
let frame_id = *frame_entry;
let frame = &self.frames[frame_id];
let mut page_guard = frame.page.write();
if let Some(page) = page_guard.as_mut() {
if page.is_dirty() {
self.file_manager.write_page(page)?;
self.stats.writes.fetch_add(1, Ordering::Relaxed);
}
}
}
Ok(())
}
pub fn flush_all(&self) -> Result<()> {
for entry in self.page_table.iter() {
let frame_id = *entry.value();
let frame = &self.frames[frame_id];
let mut page_guard = frame.page.write();
if let Some(page) = page_guard.as_mut() {
if page.is_dirty() {
self.file_manager.write_page(page)?;
self.stats.writes.fetch_add(1, Ordering::Relaxed);
}
}
}
self.file_manager.flush()?;
Ok(())
}
fn unpin_page(&self, frame_id: FrameId) {
self.frames[frame_id].unpin();
}
fn find_victim_frame(&self) -> Result<FrameId> {
let pool_size = self.frames.len();
let mut iterations = 0;
let max_iterations = pool_size * 2;
loop {
if iterations >= max_iterations {
return Err(TdbError::BufferPoolFull);
}
let hand = self.clock_hand.fetch_add(1, Ordering::AcqRel) as usize % pool_size;
let frame = &self.frames[hand];
if frame.is_pinned() {
iterations += 1;
continue;
}
if frame.is_empty() {
return Ok(hand);
}
let access_count = frame.get_access_count();
if access_count == 0 {
return Ok(hand);
}
frame.access_count.fetch_sub(1, Ordering::AcqRel);
iterations += 1;
}
}
pub fn stats(&self) -> BufferPoolStats {
self.stats.snapshot()
}
pub fn pool_size(&self) -> usize {
self.frames.len()
}
pub fn cached_pages(&self) -> usize {
self.page_table.len()
}
}
pub struct PageGuard<'a> {
frame_id: FrameId,
page_id: PageId,
buffer_pool: &'a BufferPool,
}
impl<'a> PageGuard<'a> {
pub fn page(&self) -> parking_lot::RwLockReadGuard<'_, Option<Page>> {
self.buffer_pool.frames[self.frame_id].page.read()
}
pub fn page_mut(&self) -> parking_lot::RwLockWriteGuard<'_, Option<Page>> {
self.buffer_pool.frames[self.frame_id].page.write()
}
pub fn page_id(&self) -> PageId {
self.page_id
}
}
impl<'a> Drop for PageGuard<'a> {
fn drop(&mut self) {
self.buffer_pool.unpin_page(self.frame_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
#[test]
fn test_buffer_pool_creation() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(10, fm);
assert_eq!(bp.pool_size(), 10);
assert_eq!(bp.cached_pages(), 0);
}
#[test]
fn test_buffer_pool_new_page() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(10, fm);
let guard = bp.new_page(PageType::BTreeLeaf).unwrap();
assert_eq!(guard.page_id(), 0);
assert_eq!(bp.cached_pages(), 1);
}
#[test]
fn test_buffer_pool_fetch_page() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(10, fm.clone());
let page_id = {
let guard = bp.new_page(PageType::BTreeLeaf).unwrap();
let mut page = guard.page_mut();
page.as_mut().unwrap().write_at(0, b"test data").unwrap();
guard.page_id()
};
bp.flush_page(page_id).unwrap();
let guard = bp.fetch_page(page_id).unwrap();
let page = guard.page();
let data = page.as_ref().unwrap().read_at(0, 9).unwrap();
assert_eq!(data, b"test data");
}
#[test]
fn test_buffer_pool_cache_hit() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(10, fm);
let guard1 = bp.new_page(PageType::BTreeLeaf).unwrap();
let page_id = guard1.page_id();
drop(guard1);
let _guard2 = bp.fetch_page(page_id).unwrap();
let stats = bp.stats();
assert!(stats.cache_hits.load(Ordering::Relaxed) > 0);
}
#[test]
fn test_buffer_pool_eviction() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(3, fm);
let _g1 = bp.new_page(PageType::BTreeLeaf).unwrap();
let _g2 = bp.new_page(PageType::BTreeLeaf).unwrap();
let _g3 = bp.new_page(PageType::BTreeLeaf).unwrap();
assert_eq!(bp.cached_pages(), 3);
drop(_g1);
drop(_g2);
drop(_g3);
let _g4 = bp.new_page(PageType::BTreeLeaf).unwrap();
let stats = bp.stats();
assert!(stats.evictions.load(Ordering::Relaxed) > 0);
}
#[test]
fn test_buffer_pool_flush_all() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(10, fm);
{
let g1 = bp.new_page(PageType::BTreeLeaf).unwrap();
let mut page1 = g1.page_mut();
page1.as_mut().unwrap().write_at(0, b"data1").unwrap();
}
{
let g2 = bp.new_page(PageType::BTreeLeaf).unwrap();
let mut page2 = g2.page_mut();
page2.as_mut().unwrap().write_at(0, b"data2").unwrap();
}
bp.flush_all().unwrap();
let stats = bp.stats();
assert!(stats.writes.load(Ordering::Relaxed) >= 2);
}
#[test]
fn test_buffer_pool_hit_rate() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(10, fm);
let guard = bp.new_page(PageType::BTreeLeaf).unwrap();
let page_id = guard.page_id();
drop(guard);
for _ in 0..5 {
let _g = bp.fetch_page(page_id).unwrap();
}
let stats = bp.stats();
assert!(stats.hit_rate() > 0.8); }
#[test]
fn test_buffer_pool_pin_prevents_eviction() {
let temp_file = NamedTempFile::new().unwrap();
let fm = Arc::new(FileManager::open(temp_file.path(), false).unwrap());
let bp = BufferPool::new(2, fm);
let g1 = bp.new_page(PageType::BTreeLeaf).unwrap();
let page_id1 = g1.page_id();
let g2 = bp.new_page(PageType::BTreeLeaf).unwrap();
drop(g2);
let result = bp.new_page(PageType::BTreeLeaf);
assert!(result.is_ok());
assert!(bp.page_table.contains_key(&page_id1));
}
}