use std::{mem::MaybeUninit, ptr, slice};
use wbftree::{
BfTreeDeleteResult, BfTreeInsertResult, BfTreeReadResult, BfTreeService, ScanReturnField,
};
use crate::{
CollectionError, Result,
prefix::{TreePrefix, with_prefixed_key},
};
pub const TAG_EMPTY: u8 = 0;
pub const TAG_NON_EMPTY: u8 = 1;
pub const TAG_PADDED: u8 = 2;
const STACK_VAL_BUF_SIZE: usize = 1024;
#[inline]
fn parse_hash_payload(bytes: &[u8]) -> Option<&[u8]> {
match bytes.first() {
Some(&TAG_EMPTY) => Some(&[]),
Some(&TAG_NON_EMPTY) => Some(&bytes[1..]),
Some(&TAG_PADDED) => {
if bytes.len() >= 2 {
let len = bytes[1] as usize;
if bytes.len() >= 2 + len {
return Some(&bytes[2..2 + len]);
}
}
None
}
_ => None,
}
}
#[inline]
fn with_hash_val<R>(key_len: usize, value: &[u8], f: impl FnOnce(&[u8]) -> R) -> R {
if value.is_empty() {
f(&[TAG_EMPTY, 0, 0, 0])
} else if key_len + 1 + value.len() < 4 {
let mut buf = [0u8; 4];
buf[0] = TAG_PADDED;
buf[1] = value.len() as u8;
buf[2..2 + value.len()].copy_from_slice(value);
f(&buf)
} else if value.len() < STACK_VAL_BUF_SIZE {
let mut buf = [MaybeUninit::<u8>::uninit(); STACK_VAL_BUF_SIZE];
let total_len = 1 + value.len();
unsafe {
let ptr = buf.as_mut_ptr() as *mut u8;
*ptr = TAG_NON_EMPTY;
ptr::copy_nonoverlapping(value.as_ptr(), ptr.add(1), value.len());
f(slice::from_raw_parts(buf.as_ptr() as *const u8, total_len))
}
} else {
let mut buf = Vec::with_capacity(1 + value.len());
buf.push(TAG_NON_EMPTY);
buf.extend_from_slice(value);
f(&buf)
}
}
pub trait HashTreeOps {
fn hset(&self, field: &[u8], value: &[u8]) -> Result<bool>;
fn hget(&self, field: &[u8]) -> Result<Option<Vec<u8>>> {
self.hget_callback(field, |opt| opt.map(|v| v.to_vec()))
}
fn hget_callback<R>(&self, field: &[u8], f: impl FnOnce(Option<&[u8]>) -> R) -> Result<R>;
fn hdel(&self, field: &[u8]) -> Result<bool>;
fn hexists(&self, field: &[u8]) -> Result<bool>;
fn hlen(&self) -> Result<usize>;
fn hscan<F>(&self, start_field: &[u8], count: usize, on_entry: F) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool;
}
impl HashTreeOps for BfTreeService {
fn hset(&self, field: &[u8], value: &[u8]) -> Result<bool> {
with_prefixed_key(TreePrefix::HashField.as_u8(), field, |k| {
let exists = self.contains_key(k);
let insert_res = with_hash_val(k.len(), value, |val_bytes| self.insert(k, val_bytes));
match insert_res {
BfTreeInsertResult::Success => Ok(!exists),
BfTreeInsertResult::InvalidKV => Err(CollectionError::KeyTooLong),
_ => Err(CollectionError::InvalidArgument("hset 插入失败")),
}
})
}
fn hget_callback<R>(&self, field: &[u8], f: impl FnOnce(Option<&[u8]>) -> R) -> Result<R> {
with_prefixed_key(TreePrefix::HashField.as_u8(), field, |k| {
self.read_callback(k, |res, bytes| match res {
BfTreeReadResult::Found => match parse_hash_payload(bytes) {
Some(payload) => Ok(f(Some(payload))),
None => Err(CollectionError::Corrupted("哈希值标签非法")),
},
BfTreeReadResult::NotFound | BfTreeReadResult::Deleted => Ok(f(None)),
_ => Err(CollectionError::InvalidArgument("hget 读取失败")),
})
})
}
fn hdel(&self, field: &[u8]) -> Result<bool> {
with_prefixed_key(TreePrefix::HashField.as_u8(), field, |k| {
let exists = self.contains_key(k);
if !exists {
return Ok(false);
}
match self.delete(k) {
BfTreeDeleteResult::Success => Ok(true),
_ => Err(CollectionError::InvalidArgument("hdel 删除失败")),
}
})
}
#[inline]
fn hexists(&self, field: &[u8]) -> Result<bool> {
with_prefixed_key(TreePrefix::HashField.as_u8(), field, |k| {
Ok(self.contains_key(k))
})
}
fn hlen(&self) -> Result<usize> {
let prefix_u8 = TreePrefix::HashField as u8;
let start_key = [prefix_u8];
let mut count = 0;
self.scan_with_count_callback(&start_key, usize::MAX, ScanReturnField::Key, |k, _val| {
if k.is_empty() || k[0] != prefix_u8 {
return false;
}
count += 1;
true
})?;
Ok(count)
}
fn hscan<F>(&self, start_field: &[u8], count: usize, mut on_entry: F) -> Result<usize>
where
F: FnMut(&[u8], &[u8]) -> bool,
{
if count == 0 {
return Ok(0);
}
let prefix_u8 = TreePrefix::HashField as u8;
let mut corrupted = false;
let scanned = with_prefixed_key(prefix_u8, start_field, |sk| {
self.scan_with_count_callback(sk, count, ScanReturnField::KeyAndValue, |k, v| {
if k.is_empty() || k[0] != prefix_u8 {
return false;
}
match parse_hash_payload(v) {
Some(payload) => on_entry(&k[1..], payload),
None => {
corrupted = true;
false
}
}
})
})?;
if corrupted {
return Err(CollectionError::Corrupted("哈希值标签非法"));
}
Ok(scanned)
}
}