use std::cell::RefCell;
use std::sync::Arc;
use anyhow::{Context, Result};
use kingfisher_vectorscan::{BlockDatabase, BlockScanner};
use thread_local::ThreadLocal;
pub struct ScannerPool {
scanners: ThreadLocal<RefCell<Option<BlockScanner<'static>>>>,
db: Arc<BlockDatabase>,
}
impl ScannerPool {
pub fn new(db: Arc<BlockDatabase>) -> Self {
Self { db, scanners: ThreadLocal::new() }
}
pub fn with<F, R>(&self, f: F) -> R
where
F: for<'db> FnOnce(&mut BlockScanner<'db>) -> R,
{
self.try_with(f).expect("unable to borrow or initialize scanner")
}
pub fn try_with<F, R>(&self, f: F) -> Result<R>
where
F: for<'db> FnOnce(&mut BlockScanner<'db>) -> R,
{
let cell = self.scanners.get_or(|| RefCell::new(None));
let mut scanner_opt = cell.try_borrow_mut().context("scanner pool is already borrowed")?;
if scanner_opt.is_none() {
let db_ref: &'static BlockDatabase =
unsafe { std::mem::transmute::<&BlockDatabase, &'static BlockDatabase>(&self.db) };
*scanner_opt = Some(BlockScanner::new(db_ref)?);
}
Ok(f(scanner_opt.as_mut().unwrap()))
}
}
#[cfg(test)]
mod tests {
use std::panic::{AssertUnwindSafe, catch_unwind};
use kingfisher_vectorscan::{Flag, Pattern, Scan};
use super::*;
fn database() -> Arc<BlockDatabase> {
Arc::new(
BlockDatabase::new(vec![Pattern::new(b"secret".to_vec(), Flag::default(), Some(0))])
.unwrap(),
)
}
#[test]
fn pool_retains_database_until_all_thread_scratch_is_dropped() {
let database = database();
let weak = Arc::downgrade(&database);
let pool = ScannerPool::new(database);
std::thread::scope(|scope| {
for _ in 0..4 {
let pool = &pool;
scope.spawn(move || {
let mut matches = 0;
pool.try_with(|scanner| {
scanner.scan(b"a secret", |_, _, _, _| {
matches += 1;
Scan::Continue
})
})
.unwrap()
.unwrap();
assert_eq!(matches, 1);
});
}
});
assert!(weak.upgrade().is_some());
drop(pool);
assert!(weak.upgrade().is_none());
}
#[test]
fn panicking_callback_releases_the_scanner_borrow() {
let pool = ScannerPool::new(database());
assert!(
catch_unwind(AssertUnwindSafe(|| {
let _ = pool.try_with(|_| panic!("callback failed"));
}))
.is_err()
);
pool.try_with(|_| ()).unwrap();
}
}