use std::{
cell::RefCell,
mem::MaybeUninit,
panic::{AssertUnwindSafe, catch_unwind},
result::Result as StdResult,
slice::from_raw_parts_mut,
sync::atomic::Ordering,
};
use bf_tree::{BfTree, LeafInsertResult, LeafReadResult, ScanIter, ScanIterError};
use super::{BfTreeService, MIN_MAX_RECORD_SIZE, STACK_READ_BUF_SIZE, STACK_SCAN_BUF_SIZE};
use crate::{
error::{Error, Result},
types::{BfTreeDeleteResult, BfTreeInsertResult, BfTreeReadResult, ScanRecord, ScanReturnField},
};
thread_local! {
static READ_SCRATCH: RefCell<Vec<u8>> = const { RefCell::new(Vec::new()) };
static SCAN_SCRATCH: RefCell<Vec<u8>> = const { RefCell::new(Vec::new()) };
}
#[inline]
fn with_read_scratch<R>(max_record_size: usize, f: impl FnOnce(&mut [u8]) -> R) -> R {
READ_SCRATCH.with(|scratch| {
if let Ok(mut s) = scratch.try_borrow_mut() {
if s.len() < max_record_size {
s.resize(max_record_size, 0);
}
f(&mut s[..max_record_size])
} else {
let mut fallback = vec![0u8; max_record_size];
f(&mut fallback)
}
})
}
#[inline]
fn with_scan_scratch<R>(max_record_size: usize, f: impl FnOnce(&mut [u8]) -> R) -> R {
SCAN_SCRATCH.with(|scratch| {
if let Ok(mut s) = scratch.try_borrow_mut() {
if s.len() < max_record_size {
s.resize(max_record_size, 0);
}
f(&mut s[..max_record_size])
} else {
let mut fallback = vec![0u8; max_record_size];
f(&mut fallback)
}
})
}
#[inline]
fn scan_iter_error_to_string(e: ScanIterError) -> &'static str {
match e {
ScanIterError::CacheOnlyMode => "CacheOnlyMode",
ScanIterError::InvalidStartKey => "InvalidStartKey",
ScanIterError::InvalidEndKey => "InvalidEndKey",
ScanIterError::InvalidCount => "InvalidCount",
ScanIterError::InvalidKeyRange => "InvalidKeyRange",
}
}
pub const SCAN_ALL_START_KEY: &[u8] = &[0];
impl BfTreeService {
#[inline]
pub fn insert(&self, key: &[u8], value: &[u8]) -> BfTreeInsertResult {
if value.is_empty() {
return BfTreeInsertResult::InvalidKV;
}
if self.barriers.load(Ordering::Relaxed) != 0 {
self.wait_for_barrier();
}
let Ok(tree) = self.tree_ref() else {
return BfTreeInsertResult::InvalidArguments;
};
match tree.insert(key, value) {
LeafInsertResult::Success => BfTreeInsertResult::Success,
LeafInsertResult::InvalidKV(_) => BfTreeInsertResult::InvalidKV,
}
}
pub fn read_callback<R>(&self, key: &[u8], f: impl FnOnce(BfTreeReadResult, &[u8]) -> R) -> R {
let Ok(tree) = self.tree_ref() else {
return f(BfTreeReadResult::InvalidArguments, &[]);
};
let max_record_size = self.max_record_size();
if max_record_size <= STACK_READ_BUF_SIZE {
let mut stack_buf = MaybeUninit::<[u8; STACK_READ_BUF_SIZE]>::uninit();
let buf =
unsafe { from_raw_parts_mut(stack_buf.as_mut_ptr() as *mut u8, STACK_READ_BUF_SIZE) };
let (res, len) = Self::read_tree(tree, key, buf);
f(res, &buf[..len])
} else {
with_read_scratch(max_record_size, |scratch| {
let (res, len) = Self::read_tree(tree, key, scratch);
f(res, &scratch[..len])
})
}
}
#[inline]
pub fn contains_key(&self, key: &[u8]) -> bool {
self.read_callback(key, |res, _val| res == BfTreeReadResult::Found)
}
pub fn read(&self, key: &[u8]) -> (BfTreeReadResult, Option<Vec<u8>>) {
self.read_callback(key, |res, bytes| {
(
res,
(res == BfTreeReadResult::Found).then(|| bytes.to_vec()),
)
})
}
#[inline]
fn read_tree(tree: &BfTree, key: &[u8], out_buf: &mut [u8]) -> (BfTreeReadResult, usize) {
match tree.read(key, out_buf) {
LeafReadResult::Found(n) => (BfTreeReadResult::Found, n as usize),
LeafReadResult::NotFound => (BfTreeReadResult::NotFound, 0),
LeafReadResult::Deleted => (BfTreeReadResult::Deleted, 0),
LeafReadResult::InvalidKey => (BfTreeReadResult::InvalidKey, 0),
}
}
#[inline]
fn read_into_fallback(
tree: &BfTree,
key: &[u8],
out_buf: &mut [u8],
max_record_size: usize,
) -> (BfTreeReadResult, usize) {
if max_record_size <= STACK_READ_BUF_SIZE {
let mut stack_buf = MaybeUninit::<[u8; STACK_READ_BUF_SIZE]>::uninit();
let buf =
unsafe { from_raw_parts_mut(stack_buf.as_mut_ptr() as *mut u8, STACK_READ_BUF_SIZE) };
let (res, len) = Self::read_tree(tree, key, buf);
if res == BfTreeReadResult::Found {
if len <= out_buf.len() {
out_buf[..len].copy_from_slice(&buf[..len]);
(res, len)
} else {
(BfTreeReadResult::InvalidArguments, 0)
}
} else {
(res, 0)
}
} else {
with_read_scratch(max_record_size, |scratch| {
let (res, len) = Self::read_tree(tree, key, scratch);
if res == BfTreeReadResult::Found {
if len <= out_buf.len() {
out_buf[..len].copy_from_slice(&scratch[..len]);
(res, len)
} else {
(BfTreeReadResult::InvalidArguments, 0)
}
} else {
(res, 0)
}
})
}
}
#[inline]
pub fn read_into(&self, key: &[u8], out_buf: &mut [u8]) -> (BfTreeReadResult, usize) {
let Ok(tree) = self.tree_ref() else {
return (BfTreeReadResult::InvalidArguments, 0);
};
let max_record_size = self.max_record_size();
if out_buf.len() >= max_record_size {
Self::read_tree(tree, key, out_buf)
} else {
Self::read_into_fallback(tree, key, out_buf, max_record_size)
}
}
#[inline]
pub fn delete(&self, key: &[u8]) -> BfTreeDeleteResult {
if self.barriers.load(Ordering::Relaxed) != 0 {
self.wait_for_barrier();
}
let Ok(tree) = self.tree_ref() else {
return BfTreeDeleteResult::InvalidArguments;
};
tree.delete(key);
BfTreeDeleteResult::Success
}
#[inline]
pub fn noop(&self, _key: &[u8]) -> i32 {
0
}
#[inline]
pub unsafe fn noop_by_ptr(_tree_ptr: u64, _key: &[u8]) -> i32 {
0
}
#[inline]
pub unsafe fn insert_by_ptr(tree_ptr: u64, key: &[u8], value: &[u8]) -> BfTreeInsertResult {
if tree_ptr == 0 {
return BfTreeInsertResult::InvalidArguments;
}
if value.is_empty() {
return BfTreeInsertResult::InvalidKV;
}
let tree = unsafe { &*(tree_ptr as usize as *const BfTree) };
match tree.insert(key, value) {
LeafInsertResult::Success => BfTreeInsertResult::Success,
LeafInsertResult::InvalidKV(_) => BfTreeInsertResult::InvalidKV,
}
}
#[inline]
pub unsafe fn read_by_ptr_into(
tree_ptr: u64,
key: &[u8],
out_buf: &mut [u8],
) -> (BfTreeReadResult, usize) {
if tree_ptr == 0 {
return (BfTreeReadResult::InvalidArguments, 0);
}
let tree = unsafe { &*(tree_ptr as usize as *const BfTree) };
let max_record_size = tree
.config()
.get_cb_max_record_size()
.max(MIN_MAX_RECORD_SIZE);
if out_buf.len() >= max_record_size {
Self::read_tree(tree, key, out_buf)
} else {
Self::read_into_fallback(tree, key, out_buf, max_record_size)
}
}
#[inline]
pub unsafe fn read_by_ptr(tree_ptr: u64, key: &[u8]) -> (BfTreeReadResult, Option<Vec<u8>>) {
if tree_ptr == 0 {
return (BfTreeReadResult::InvalidArguments, None);
}
let tree = unsafe { &*(tree_ptr as usize as *const BfTree) };
let max_record_size = tree
.config()
.get_cb_max_record_size()
.max(MIN_MAX_RECORD_SIZE);
if max_record_size <= STACK_READ_BUF_SIZE {
let mut stack_buf = MaybeUninit::<[u8; STACK_READ_BUF_SIZE]>::uninit();
let buf =
unsafe { from_raw_parts_mut(stack_buf.as_mut_ptr() as *mut u8, STACK_READ_BUF_SIZE) };
let (res, len) = Self::read_tree(tree, key, buf);
(
res,
(res == BfTreeReadResult::Found).then(|| buf[..len].to_vec()),
)
} else {
with_read_scratch(max_record_size, |scratch| {
let (res, len) = Self::read_tree(tree, key, scratch);
(
res,
(res == BfTreeReadResult::Found).then(|| scratch[..len].to_vec()),
)
})
}
}
#[inline]
pub unsafe fn delete_by_ptr(tree_ptr: u64, key: &[u8]) -> BfTreeDeleteResult {
if tree_ptr == 0 {
return BfTreeDeleteResult::InvalidArguments;
}
let tree = unsafe { &*(tree_ptr as usize as *const BfTree) };
tree.delete(key);
BfTreeDeleteResult::Success
}
pub unsafe fn scan_with_count_by_ptr_callback<F>(
tree_ptr: u64,
start_key: &[u8],
count: usize,
return_field: ScanReturnField,
on_record: F,
) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
if tree_ptr == 0 {
return Err(Error::InvalidArgument("原生树指针为空".into()));
}
if count == 0 {
return Ok(0);
}
let tree = unsafe { &*(tree_ptr as usize as *const BfTree) };
let mut iter = match catch_unwind(AssertUnwindSafe(|| {
tree.scan_with_count(start_key, count, return_field.into())
})) {
Ok(Ok(iter)) => iter,
Ok(Err(e)) => {
return Err(Error::InvalidArgument(
scan_iter_error_to_string(e).to_string(),
));
}
Err(_) => return Err(Error::Scan("底层引擎扫描初始化异常".into())),
};
let max_record_size = tree
.config()
.get_cb_max_record_size()
.max(MIN_MAX_RECORD_SIZE);
if max_record_size <= STACK_SCAN_BUF_SIZE {
let mut stack_buf = MaybeUninit::<[u8; STACK_SCAN_BUF_SIZE]>::uninit();
let buf =
unsafe { from_raw_parts_mut(stack_buf.as_mut_ptr() as *mut u8, STACK_SCAN_BUF_SIZE) };
catch_unwind(AssertUnwindSafe(|| {
Self::drain_scan_iter(&mut iter, buf, return_field, on_record)
}))
.unwrap_or_else(|_| Err(Error::Scan("底层引擎扫描排空异常".into())))
} else {
with_scan_scratch(max_record_size, |buf| {
catch_unwind(AssertUnwindSafe(|| {
Self::drain_scan_iter(&mut iter, buf, return_field, on_record)
}))
.unwrap_or_else(|_| Err(Error::Scan("底层引擎扫描排空异常".into())))
})
}
}
pub unsafe fn scan_with_end_key_by_ptr_callback<F>(
tree_ptr: u64,
start_key: &[u8],
end_key: &[u8],
return_field: ScanReturnField,
on_record: F,
) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
if tree_ptr == 0 {
return Err(Error::InvalidArgument("原生树指针为空".into()));
}
if start_key > end_key {
return Ok(0);
}
let tree = unsafe { &*(tree_ptr as usize as *const BfTree) };
let mut iter = match catch_unwind(AssertUnwindSafe(|| {
tree.scan_with_end_key(start_key, end_key, return_field.into())
})) {
Ok(Ok(iter)) => iter,
Ok(Err(e)) => {
return Err(Error::InvalidArgument(
scan_iter_error_to_string(e).to_string(),
));
}
Err(_) => return Err(Error::Scan("底层引擎扫描初始化异常".into())),
};
let max_record_size = tree
.config()
.get_cb_max_record_size()
.max(MIN_MAX_RECORD_SIZE);
if max_record_size <= STACK_SCAN_BUF_SIZE {
let mut stack_buf = MaybeUninit::<[u8; STACK_SCAN_BUF_SIZE]>::uninit();
let buf =
unsafe { from_raw_parts_mut(stack_buf.as_mut_ptr() as *mut u8, STACK_SCAN_BUF_SIZE) };
catch_unwind(AssertUnwindSafe(|| {
Self::drain_scan_iter(&mut iter, buf, return_field, on_record)
}))
.unwrap_or_else(|_| Err(Error::Scan("底层引擎扫描排空异常".into())))
} else {
with_scan_scratch(max_record_size, |buf| {
catch_unwind(AssertUnwindSafe(|| {
Self::drain_scan_iter(&mut iter, buf, return_field, on_record)
}))
.unwrap_or_else(|_| Err(Error::Scan("底层引擎扫描排空异常".into())))
})
}
}
pub fn scan_with_count_callback<F>(
&self,
start_key: &[u8],
count: usize,
return_field: ScanReturnField,
on_record: F,
) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
if count == 0 {
return Ok(0);
}
self.scan_callback(
|tree| tree.scan_with_count(start_key, count, return_field.into()),
return_field,
on_record,
)
}
pub fn scan_with_count(
&self,
start_key: &[u8],
count: usize,
return_field: ScanReturnField,
) -> Result<Vec<ScanRecord>> {
let mut records = Vec::with_capacity(count.min(64));
self.scan_with_count_callback(
start_key,
count,
return_field,
ScanRecord::sink(&mut records),
)?;
Ok(records)
}
pub fn scan_with_end_key_callback<F>(
&self,
start_key: &[u8],
end_key: &[u8],
return_field: ScanReturnField,
on_record: F,
) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
if start_key > end_key {
return Ok(0);
}
self.scan_callback(
|tree| tree.scan_with_end_key(start_key, end_key, return_field.into()),
return_field,
on_record,
)
}
pub fn scan_with_end_key(
&self,
start_key: &[u8],
end_key: &[u8],
return_field: ScanReturnField,
) -> Result<Vec<ScanRecord>> {
let mut records = Vec::with_capacity(32);
self.scan_with_end_key_callback(
start_key,
end_key,
return_field,
ScanRecord::sink(&mut records),
)?;
Ok(records)
}
pub fn scan_all_callback<F>(&self, return_field: ScanReturnField, on_record: F) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
self.scan_with_count_callback(SCAN_ALL_START_KEY, usize::MAX, return_field, on_record)
}
pub fn scan_all(&self, return_field: ScanReturnField) -> Result<Vec<ScanRecord>> {
let mut records = Vec::with_capacity(32);
self.scan_all_callback(return_field, ScanRecord::sink(&mut records))?;
Ok(records)
}
#[inline]
fn drain_scan_iter<F>(
iter: &mut ScanIter<'_, '_>,
buf: &mut [u8],
return_field: ScanReturnField,
mut on_record: F,
) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
let mut scanned = 0;
match return_field {
ScanReturnField::KeyAndValue => {
while let Some((k_len, v_len)) = iter.next(buf) {
scanned += 1;
if !on_record(&buf[..k_len], &buf[k_len..k_len + v_len]) {
break;
}
}
}
ScanReturnField::Key => {
while let Some((k_len, _)) = iter.next(buf) {
scanned += 1;
if !on_record(&buf[..k_len], &[]) {
break;
}
}
}
ScanReturnField::Value => {
while let Some((k_len, v_len)) = iter.next(buf) {
scanned += 1;
if !on_record(&[], &buf[k_len..k_len + v_len]) {
break;
}
}
}
}
Ok(scanned)
}
fn scan_callback<F>(
&self,
make_iter: impl FnOnce(&BfTree) -> StdResult<ScanIter<'_, '_>, ScanIterError>,
return_field: ScanReturnField,
on_record: F,
) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
let tree = self.tree_arc()?;
let mut iter = match catch_unwind(AssertUnwindSafe(|| make_iter(&tree))) {
Ok(Ok(iter)) => iter,
Ok(Err(e)) => {
return Err(Error::InvalidArgument(
scan_iter_error_to_string(e).to_string(),
));
}
Err(_) => return Err(Error::Scan("底层引擎扫描初始化异常".into())),
};
let max_record_size = self.max_record_size();
if max_record_size <= STACK_SCAN_BUF_SIZE {
let mut stack_buf = MaybeUninit::<[u8; STACK_SCAN_BUF_SIZE]>::uninit();
let buf =
unsafe { from_raw_parts_mut(stack_buf.as_mut_ptr() as *mut u8, STACK_SCAN_BUF_SIZE) };
catch_unwind(AssertUnwindSafe(|| {
Self::drain_scan_iter(&mut iter, buf, return_field, on_record)
}))
.unwrap_or_else(|_| Err(Error::Scan("底层引擎扫描排空异常".into())))
} else {
with_scan_scratch(max_record_size, |buf| {
catch_unwind(AssertUnwindSafe(|| {
Self::drain_scan_iter(&mut iter, buf, return_field, on_record)
}))
.unwrap_or_else(|_| Err(Error::Scan("底层引擎扫描排空异常".into())))
})
}
}
}