#![allow(unsafe_code)]
use std::marker::PhantomData;
use std::ptr::NonNull;
use super::arena::Arena;
use super::internal_key::{INTERNAL_KEY_SUFFIX_LEN, compare_internal_keys, compare_internal_split};
use crate::sync::{Arc, AtomicPtr, AtomicU64, AtomicUsize, Ordering};
const MAX_HEIGHT: usize = 12;
const BRANCHING: u64 = 4;
const NODE_HEADER: usize = 16;
type Link = AtomicPtr<u8>;
const PTR_SIZE: usize = size_of::<Link>();
pub(crate) const NODE_ALIGN: usize = if align_of::<Link>() > 8 {
align_of::<Link>()
} else {
8
};
const HEAD_SIZE: usize = NODE_HEADER + PTR_SIZE * MAX_HEIGHT;
fn node_size(key_len: usize, value_len: usize, height: usize) -> usize {
(NODE_HEADER + PTR_SIZE * height + key_len + value_len).next_multiple_of(NODE_ALIGN)
}
pub(crate) fn max_node_size(internal_key_len: usize, value_len: usize) -> usize {
node_size(internal_key_len, value_len, MAX_HEIGHT)
}
unsafe fn header(node: *const u8) -> (usize, usize, usize) {
unsafe {
let key_len = node.cast::<u32>().read() as usize;
let value_len = node.add(4).cast::<u32>().read() as usize;
let height = node.add(8).read() as usize;
(key_len, value_len, height)
}
}
unsafe fn link(node: *const u8, level: usize) -> &'static Link {
unsafe { &*node.add(NODE_HEADER + level * PTR_SIZE).cast::<Link>() }
}
unsafe fn next_at(node: *const u8, level: usize) -> Option<NonNull<u8>> {
NonNull::new(unsafe { link(node, level) }.load(Ordering::Acquire))
}
unsafe fn node_key<'a>(node: *const u8) -> &'a [u8] {
unsafe {
let (key_len, _, height) = header(node);
std::slice::from_raw_parts(node.add(NODE_HEADER + PTR_SIZE * height), key_len)
}
}
#[derive(Clone, Copy)]
pub(crate) struct NodeRef<'a> {
ptr: NonNull<u8>,
_marker: PhantomData<&'a ArenaSkipList>,
}
impl<'a> NodeRef<'a> {
fn new(ptr: NonNull<u8>) -> Self {
Self {
ptr,
_marker: PhantomData,
}
}
pub(crate) fn key(&self) -> &'a [u8] {
unsafe { node_key(self.ptr.as_ptr()) }
}
pub(crate) fn value(&self) -> &'a [u8] {
let (ptr, len) = self.value_span();
match ptr {
Some(ptr) => unsafe { std::slice::from_raw_parts(ptr.as_ptr(), len) },
None => &[],
}
}
pub(crate) fn key_span(&self) -> (Option<NonNull<u8>>, usize) {
unsafe {
let node = self.ptr.as_ptr();
let (key_len, _, height) = header(node);
if key_len == 0 {
return (None, 0);
}
let ptr = node.add(NODE_HEADER + PTR_SIZE * height);
(NonNull::new(ptr), key_len)
}
}
pub(crate) fn value_span(&self) -> (Option<NonNull<u8>>, usize) {
unsafe {
let node = self.ptr.as_ptr();
let (key_len, value_len, height) = header(node);
if value_len == 0 {
return (None, 0);
}
let ptr = node.add(NODE_HEADER + PTR_SIZE * height + key_len);
(NonNull::new(ptr), value_len)
}
}
pub(crate) fn next(&self) -> Option<NodeRef<'a>> {
unsafe { next_at(self.ptr.as_ptr(), 0) }.map(NodeRef::new)
}
}
pub(crate) struct ArenaSkipList {
arena: Arc<Arena>,
head: NonNull<u8>,
rnd: AtomicU64,
count: AtomicUsize,
#[cfg(debug_assertions)]
inserting: crate::sync::AtomicBool,
}
unsafe impl Send for ArenaSkipList {}
unsafe impl Sync for ArenaSkipList {}
impl ArenaSkipList {
pub(crate) fn new(arena: Arc<Arena>) -> Option<Self> {
let layout = std::alloc::Layout::from_size_align(HEAD_SIZE, NODE_ALIGN).ok()?;
let head = NonNull::new(unsafe { std::alloc::alloc(layout) })?;
unsafe { init_node_header(head.as_ptr(), 0, 0, MAX_HEIGHT) };
Some(Self {
arena,
head,
rnd: AtomicU64::new(0x2545_F491_4F6C_DD1D),
count: AtomicUsize::new(0),
#[cfg(debug_assertions)]
inserting: crate::sync::AtomicBool::new(false),
})
}
pub(crate) fn arena(&self) -> &Arc<Arena> {
&self.arena
}
pub(crate) fn is_empty(&self) -> bool {
self.first().is_none()
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.count.load(Ordering::Acquire)
}
fn random_height(&self) -> usize {
let mut x = self.rnd.load(Ordering::Relaxed);
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.rnd.store(x, Ordering::Relaxed);
let mut height = 1;
let mut draw = x;
while height < MAX_HEIGHT && draw.is_multiple_of(BRANCHING) {
height += 1;
draw /= BRANCHING;
}
height
}
pub(crate) fn insert(&self, user_key: &[u8], seq: u64, value_type: u8, value: &[u8]) -> bool {
#[cfg(debug_assertions)]
let _writer = SingleWriterGuard::enter(&self.inserting);
let key_len = user_key.len() + INTERNAL_KEY_SUFFIX_LEN;
if u32::try_from(key_len).is_err() || u32::try_from(value.len()).is_err() {
tracing::error!(
key_len,
value_len = value.len(),
"memtable entry too large to encode; write dropped"
);
return false;
}
let trailer = internal_trailer(seq, value_type);
let height = self.random_height();
let mut prev = [self.head; MAX_HEIGHT];
let mut cursor = self.head;
for level in (0..MAX_HEIGHT).rev() {
while let Some(next) = unsafe { next_at(cursor.as_ptr(), level) } {
let next_key = unsafe { node_key(next.as_ptr()) };
if compare_internal_split(next_key, user_key, &trailer).is_lt() {
cursor = next;
} else {
break;
}
}
prev[level] = cursor;
}
let size = node_size(key_len, value.len(), height);
let Some(node) = self.arena.alloc(size, NODE_ALIGN) else {
let layout = std::alloc::Layout::from_size_align(size, NODE_ALIGN)
.unwrap_or_else(|_| std::alloc::Layout::new::<u8>());
std::alloc::handle_alloc_error(layout)
};
unsafe {
let raw = node.as_ptr();
init_node_header(raw, key_len, value.len(), height);
let key_at = raw.add(NODE_HEADER + PTR_SIZE * height);
std::ptr::copy_nonoverlapping(user_key.as_ptr(), key_at, user_key.len());
std::ptr::copy_nonoverlapping(
trailer.as_ptr(),
key_at.add(user_key.len()),
INTERNAL_KEY_SUFFIX_LEN,
);
std::ptr::copy_nonoverlapping(value.as_ptr(), key_at.add(key_len), value.len());
for (level, slot) in prev.iter().enumerate().take(height) {
let successor =
next_at(slot.as_ptr(), level).map_or(std::ptr::null_mut(), |n| n.as_ptr());
link(raw, level).store(successor, Ordering::Relaxed);
}
for (level, slot) in prev.iter().enumerate().take(height) {
link(slot.as_ptr(), level).store(raw, Ordering::Release);
}
}
self.count.fetch_add(1, Ordering::Release);
true
}
fn seek_pair(&self, target: &[u8]) -> (Option<NonNull<u8>>, Option<NonNull<u8>>) {
let mut cursor = self.head;
let mut successor = None;
for level in (0..MAX_HEIGHT).rev() {
loop {
let next = unsafe { next_at(cursor.as_ptr(), level) };
match next {
Some(node)
if compare_internal_keys(unsafe { node_key(node.as_ptr()) }, target)
.is_lt() =>
{
cursor = node;
}
other => {
if level == 0 {
successor = other;
}
break;
}
}
}
}
((cursor != self.head).then_some(cursor), successor)
}
pub(crate) fn seek_ge(&self, target: &[u8]) -> Option<NodeRef<'_>> {
self.seek_pair(target).1.map(NodeRef::new)
}
pub(crate) fn seek_gt(&self, target: &[u8]) -> Option<NodeRef<'_>> {
let mut node = self.seek_ge(target);
while let Some(current) = node {
if compare_internal_keys(current.key(), target).is_gt() {
return Some(current);
}
node = current.next();
}
None
}
pub(crate) fn seek_le(&self, target: &[u8]) -> Option<NodeRef<'_>> {
let (below, at_or_after) = self.seek_pair(target);
if let Some(node) = at_or_after {
let node = NodeRef::new(node);
if compare_internal_keys(node.key(), target).is_le() {
return Some(node);
}
}
below.map(NodeRef::new)
}
pub(crate) fn seek_lt(&self, target: &[u8]) -> Option<NodeRef<'_>> {
self.seek_pair(target).0.map(NodeRef::new)
}
pub(crate) fn first(&self) -> Option<NodeRef<'_>> {
unsafe { next_at(self.head.as_ptr(), 0) }.map(NodeRef::new)
}
pub(crate) fn last(&self) -> Option<NodeRef<'_>> {
let mut cursor = self.head;
for level in (0..MAX_HEIGHT).rev() {
while let Some(next) = unsafe { next_at(cursor.as_ptr(), level) } {
cursor = next;
}
}
(cursor != self.head).then(|| NodeRef::new(cursor))
}
}
impl Drop for ArenaSkipList {
fn drop(&mut self) {
unsafe {
let layout = std::alloc::Layout::from_size_align_unchecked(HEAD_SIZE, NODE_ALIGN);
std::alloc::dealloc(self.head.as_ptr(), layout);
}
}
}
fn internal_trailer(seq: u64, value_type: u8) -> [u8; INTERNAL_KEY_SUFFIX_LEN] {
let mut trailer = [0u8; INTERNAL_KEY_SUFFIX_LEN];
trailer[..8].copy_from_slice(&(!seq).to_be_bytes());
trailer[8] = value_type;
trailer
}
unsafe fn init_node_header(node: *mut u8, key_len: usize, value_len: usize, height: usize) {
debug_assert!((1..=MAX_HEIGHT).contains(&height));
unsafe {
node.cast::<u32>().write(key_len as u32);
node.add(4).cast::<u32>().write(value_len as u32);
node.add(8).write(height as u8);
std::ptr::write_bytes(node.add(9), 0, NODE_HEADER - 9);
for level in 0..height {
node.add(NODE_HEADER + level * PTR_SIZE)
.cast::<Link>()
.write(Link::new(std::ptr::null_mut()));
}
}
}
#[cfg(debug_assertions)]
struct SingleWriterGuard<'a>(&'a crate::sync::AtomicBool);
#[cfg(debug_assertions)]
impl<'a> SingleWriterGuard<'a> {
fn enter(flag: &'a crate::sync::AtomicBool) -> Self {
let busy = flag.swap(true, Ordering::Acquire);
debug_assert!(
!busy,
"ArenaSkipList::insert is single-writer (S2); the engine must serialize writers"
);
Self(flag)
}
}
#[cfg(debug_assertions)]
impl Drop for SingleWriterGuard<'_> {
fn drop(&mut self) {
self.0.store(false, Ordering::Release);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::engine::arena::{ArenaProfile, ChunkPool};
use crate::engine::internal_key::{VALUE_TYPE_DELETION, VALUE_TYPE_VALUE, encode_internal_key};
use proptest::prelude::*;
fn list(budget: usize) -> ArenaSkipList {
let profile = ArenaProfile::EMBEDDED;
let pool = Arc::new(ChunkPool::new(profile, budget, 2));
let arena = Arc::new(Arena::new(pool, budget, profile));
ArenaSkipList::new(arena).expect("head allocation")
}
fn collect(list: &ArenaSkipList) -> Vec<(Vec<u8>, Vec<u8>)> {
let mut out = Vec::new();
let mut node = list.first();
while let Some(current) = node {
out.push((current.key().to_vec(), current.value().to_vec()));
node = current.next();
}
out
}
#[test]
fn empty_list_has_no_entries() {
let list = list(64 * 1024);
assert!(list.is_empty());
assert_eq!(list.len(), 0);
assert!(list.first().is_none());
assert!(list.last().is_none());
assert!(list.seek_ge(b"anything").is_none());
assert!(list.seek_le(b"anything").is_none());
assert_eq!(
list.arena().reserved_bytes(),
0,
"head is not an arena chunk"
);
}
#[test]
fn insert_then_read_back() {
let list = list(64 * 1024);
assert!(list.insert(b"k", 1, VALUE_TYPE_VALUE, b"v1"));
assert!(list.insert(b"k", 2, VALUE_TYPE_VALUE, b"v2"));
assert!(list.insert(b"a", 3, VALUE_TYPE_DELETION, b""));
assert_eq!(list.len(), 3);
let entries = collect(&list);
assert_eq!(entries.len(), 3);
assert_eq!(
entries[0].0,
encode_internal_key(b"a", 3, VALUE_TYPE_DELETION)
);
assert_eq!(entries[0].1, b"");
assert_eq!(entries[1].0, encode_internal_key(b"k", 2, VALUE_TYPE_VALUE));
assert_eq!(entries[1].1, b"v2");
assert_eq!(entries[2].1, b"v1");
}
#[test]
fn seeks_bracket_the_list() {
let list = list(64 * 1024);
for key in [b"b".as_slice(), b"m", b"y"] {
assert!(list.insert(key, 1, VALUE_TYPE_VALUE, key));
}
let probe = |k: &[u8]| encode_internal_key(k, u64::MAX, VALUE_TYPE_DELETION);
assert_eq!(list.seek_ge(&probe(b"a")).expect("first").value(), b"b");
assert_eq!(list.seek_ge(&probe(b"m")).expect("exact").value(), b"m");
assert_eq!(list.seek_ge(&probe(b"n")).expect("next").value(), b"y");
assert!(list.seek_ge(&probe(b"z")).is_none());
assert!(list.seek_lt(&probe(b"b")).is_none());
assert_eq!(list.seek_lt(&probe(b"n")).expect("prev").value(), b"m");
assert_eq!(list.last().expect("last").value(), b"y");
assert_eq!(list.first().expect("first").value(), b"b");
}
#[test]
fn duplicate_internal_key_finds_the_first() {
let list = list(64 * 1024);
assert!(list.insert(b"k", 7, VALUE_TYPE_VALUE, b"first"));
assert!(list.insert(b"k", 7, VALUE_TYPE_VALUE, b"second"));
assert_eq!(list.len(), 2);
let probe = encode_internal_key(b"k", 7, VALUE_TYPE_DELETION);
let found = list.seek_ge(&probe).expect("present");
assert!(found.value() == b"first" || found.value() == b"second");
assert_eq!(collect(&list).len(), 2);
}
#[test]
fn prefix_keys_sort_by_the_internal_comparator() {
let list = list(64 * 1024);
assert!(list.insert(b"abc", 1, VALUE_TYPE_VALUE, b"abc"));
assert!(list.insert(b"ab", u64::MAX, VALUE_TYPE_VALUE, b"ab-high"));
assert!(list.insert(b"ab", 0, VALUE_TYPE_VALUE, b"ab-low"));
let values: Vec<Vec<u8>> = collect(&list).into_iter().map(|(_, v)| v).collect();
assert_eq!(
values,
vec![b"ab-high".to_vec(), b"ab-low".to_vec(), b"abc".to_vec()]
);
}
#[test]
fn empty_value_round_trips() {
let list = list(64 * 1024);
assert!(list.insert(b"k", 1, VALUE_TYPE_DELETION, b""));
let node = list.first().expect("present");
assert_eq!(node.value(), b"");
assert_eq!(node.value_span().1, 0);
assert!(node.value_span().0.is_none());
}
#[test]
fn a_value_larger_than_a_chunk_still_lands() {
let list = list(64 * 1024);
let big = vec![7u8; 300 * 1024];
assert!(list.insert(b"big", 1, VALUE_TYPE_VALUE, &big));
assert_eq!(list.first().expect("present").value(), big.as_slice());
}
#[test]
#[cfg_attr(
miri,
ignore = "20k inserts against 8 spinning readers; \
`a_reader_walks_the_list_while_a_writer_publishes` is the miri-sized form"
)]
fn readers_see_whole_entries_while_a_writer_inserts() {
let budget = 4 * 1024 * 1024;
let profile = ArenaProfile::SERVER;
let pool = Arc::new(ChunkPool::new(profile, budget, 2));
let arena = Arc::new(Arena::new(pool, budget, profile));
let list = Arc::new(ArenaSkipList::new(arena).expect("head"));
let stop = Arc::new(std::sync::atomic::AtomicBool::new(false));
let readers: Vec<_> = (0..8)
.map(|_| {
let list = Arc::clone(&list);
let stop = Arc::clone(&stop);
std::thread::spawn(move || {
while !stop.load(Ordering::Relaxed) {
let mut node = list.first();
while let Some(current) = node {
let key = current.key();
assert!(key.len() >= INTERNAL_KEY_SUFFIX_LEN);
let user = &key[..key.len() - INTERNAL_KEY_SUFFIX_LEN];
assert_eq!(
current.value(),
user,
"value must match the key it was written with"
);
node = current.next();
}
}
})
})
.collect();
for i in 0..20_000u32 {
let key = format!("key{i:08}");
assert!(list.insert(
key.as_bytes(),
u64::from(i) + 1,
VALUE_TYPE_VALUE,
key.as_bytes()
));
}
stop.store(true, Ordering::Relaxed);
for reader in readers {
reader.join().expect("reader thread");
}
assert_eq!(list.len(), 20_000);
}
#[test]
#[cfg_attr(
miri,
ignore = "40k inserts against 4 spinning readers; \
`a_reader_walks_the_list_while_a_writer_publishes` is the miri-sized form"
)]
fn a_seeded_key_stays_findable_while_a_writer_inserts_around_it() {
let budget = 8 * 1024 * 1024;
let profile = ArenaProfile::SERVER;
let pool = Arc::new(ChunkPool::new(profile, budget, 2));
let arena = Arc::new(Arena::new(pool, budget, profile));
let list = Arc::new(ArenaSkipList::new(arena).expect("head"));
assert!(list.insert(b"pinned", 1, VALUE_TYPE_VALUE, b"stable"));
let probe = encode_internal_key(b"pinned", u64::MAX, VALUE_TYPE_DELETION);
let stop = Arc::new(std::sync::atomic::AtomicBool::new(false));
let misses = Arc::new(AtomicUsize::new(0));
let readers: Vec<_> = (0..4)
.map(|_| {
let list = Arc::clone(&list);
let stop = Arc::clone(&stop);
let misses = Arc::clone(&misses);
let probe = probe.clone();
std::thread::spawn(move || {
while !stop.load(Ordering::Relaxed) {
let found = list
.seek_ge(&probe)
.filter(|entry| entry.value() == b"stable");
if found.is_none() {
misses.fetch_add(1, Ordering::Relaxed);
}
}
})
})
.collect();
for i in 0..40_000u64 {
let key = if i % 2 == 0 {
format!("pinne{i:08}")
} else {
format!("pinnee{i:08}")
};
assert!(list.insert(key.as_bytes(), i + 2, VALUE_TYPE_VALUE, b"noise"));
}
stop.store(true, Ordering::Relaxed);
for reader in readers {
reader.join().expect("reader thread");
}
assert_eq!(
misses.load(Ordering::Relaxed),
0,
"a present key must never read as absent"
);
}
#[test]
fn a_reader_walks_the_list_while_a_writer_publishes() {
let skiplist = Arc::new(list(64 * 1024));
assert!(skiplist.insert(b"seed", 1, VALUE_TYPE_VALUE, b"seed"));
let probe = encode_internal_key(b"seed", u64::MAX, VALUE_TYPE_DELETION);
let reader = {
let skiplist = Arc::clone(&skiplist);
let probe = probe.clone();
std::thread::spawn(move || {
for _ in 0..8 {
let mut cursor = skiplist.first();
while let Some(current) = cursor {
let key = current.key();
assert!(key.len() >= INTERNAL_KEY_SUFFIX_LEN);
let user = &key[..key.len() - INTERNAL_KEY_SUFFIX_LEN];
assert_eq!(current.value(), user);
cursor = current.next();
}
assert!(
skiplist.seek_ge(&probe).is_some(),
"S3: a published key cannot be lost"
);
}
})
};
for i in 0..8u64 {
let key = format!("k{i}");
assert!(skiplist.insert(key.as_bytes(), i + 2, VALUE_TYPE_VALUE, key.as_bytes()));
}
reader.join().expect("reader thread");
assert_eq!(skiplist.len(), 9);
}
proptest! {
#[test]
fn insert_then_seek_round_trips(
entries in proptest::collection::vec(
(proptest::collection::vec(any::<u8>(), 0..24), 1u64..1000, any::<Vec<u8>>()),
1..80,
),
) {
let list = list(1024 * 1024);
let mut model: Vec<(Vec<u8>, Vec<u8>)> = Vec::new();
for (key, seq, value) in &entries {
prop_assert!(list.insert(key, *seq, VALUE_TYPE_VALUE, value));
model.push((encode_internal_key(key, *seq, VALUE_TYPE_VALUE), value.clone()));
}
model.sort_by(|a, b| compare_internal_keys(&a.0, &b.0));
let got = collect(&list);
prop_assert_eq!(got.len(), model.len());
for (got, want) in got.iter().zip(model.iter()) {
prop_assert_eq!(&got.0, &want.0);
}
let mut got_pairs = got.clone();
let mut want_pairs = model.clone();
got_pairs.sort();
want_pairs.sort();
prop_assert_eq!(got_pairs, want_pairs);
for (key, _) in &model {
let found = list.seek_ge(key).expect("stored key is findable");
prop_assert!(compare_internal_keys(found.key(), key).is_eq());
}
}
#[test]
fn seeks_agree_with_a_linear_scan(
keys in proptest::collection::vec(proptest::collection::vec(any::<u8>(), 0..12), 1..40),
probe in proptest::collection::vec(any::<u8>(), 0..12),
) {
let list = list(1024 * 1024);
for (i, key) in keys.iter().enumerate() {
prop_assert!(list.insert(key, i as u64 + 1, VALUE_TYPE_VALUE, key));
}
let sorted = collect(&list);
let target = encode_internal_key(&probe, u64::MAX, VALUE_TYPE_DELETION);
let want_ge = sorted.iter().find(|(k, _)| compare_internal_keys(k, &target).is_ge());
let got_ge = list.seek_ge(&target).map(|n| n.key().to_vec());
prop_assert_eq!(got_ge.as_deref(), want_ge.map(|(k, _)| k.as_slice()));
let want_lt = sorted.iter().rev().find(|(k, _)| compare_internal_keys(k, &target).is_lt());
let got_lt = list.seek_lt(&target).map(|n| n.key().to_vec());
prop_assert_eq!(got_lt.as_deref(), want_lt.map(|(k, _)| k.as_slice()));
let want_le = sorted.iter().rev().find(|(k, _)| compare_internal_keys(k, &target).is_le());
let got_le = list.seek_le(&target).map(|n| n.key().to_vec());
prop_assert_eq!(got_le.as_deref(), want_le.map(|(k, _)| k.as_slice()));
let want_gt = sorted.iter().find(|(k, _)| compare_internal_keys(k, &target).is_gt());
let got_gt = list.seek_gt(&target).map(|n| n.key().to_vec());
prop_assert_eq!(got_gt.as_deref(), want_gt.map(|(k, _)| k.as_slice()));
}
#[test]
fn node_accounting_matches_the_layout(
key in proptest::collection::vec(any::<u8>(), 0..2048),
value in proptest::collection::vec(any::<u8>(), 0..4096),
) {
let list = list(1024 * 1024);
prop_assert!(list.insert(&key, 5, VALUE_TYPE_VALUE, &value));
let node = list.first().expect("present");
prop_assert_eq!(node.key().len(), key.len() + INTERNAL_KEY_SUFFIX_LEN);
prop_assert_eq!(node.value(), value.as_slice());
let used = list.arena().used_bytes();
prop_assert!(used >= key.len() + value.len() + INTERNAL_KEY_SUFFIX_LEN + NODE_HEADER);
}
}
}