use super::types::UID;
#[cfg(not(target_env = "msvc"))]
mod layout {
pub const HASH_MULTIPLIER: usize = 4;
pub const KEY_VALS: usize = 12;
pub const INFO: usize = 16;
pub const NUM_ELEMENTS: usize = 20;
pub const MASK: usize = 24;
pub const INFO_INC: usize = 32;
pub const INFO_HASH_SHIFT: usize = 36;
}
#[cfg(target_env = "msvc")]
mod layout {
pub const HASH_MULTIPLIER: usize = 16;
pub const KEY_VALS: usize = 24;
pub const INFO: usize = 28;
pub const NUM_ELEMENTS: usize = 32;
pub const MASK: usize = 36;
pub const INFO_INC: usize = 44;
pub const INFO_HASH_SHIFT: usize = 48;
}
const NODE_SIZE: usize = 16;
const NODE_VALUE_OFFSET: usize = 8;
const INFO_BITS: u32 = 5;
const INFO_MASK: u64 = (1 << INFO_BITS) - 1;
const SLOT_GET_EXTENSION_ID: usize = 0;
fn hash_int(mut x: u64) -> u64 {
x ^= x >> 33;
x = x.wrapping_mul(0xff51_afd7_ed55_8ccd);
x ^= x >> 33;
u64::from(x as u32)
}
fn hash_uid(uid: UID) -> u64 {
hash_int(uid)
}
unsafe fn read_usize(base: *mut u8, offset: usize) -> usize {
unsafe { base.add(offset).cast::<usize>().read_unaligned() }
}
#[must_use]
pub unsafe fn extension(extensible: *mut u8, uid: UID) -> *mut u8 {
if extensible.is_null() {
return std::ptr::null_mut();
}
let key_vals = unsafe { read_usize(extensible, layout::KEY_VALS) } as *mut u8;
let info = unsafe { read_usize(extensible, layout::INFO) } as *const u8;
let elements = unsafe { read_usize(extensible, layout::NUM_ELEMENTS) };
let mask = unsafe { read_usize(extensible, layout::MASK) };
if key_vals.is_null() || info.is_null() || elements == 0 {
return std::ptr::null_mut();
}
let multiplier = unsafe {
extensible
.add(layout::HASH_MULTIPLIER)
.cast::<u64>()
.read_unaligned()
};
let info_inc = unsafe {
extensible
.add(layout::INFO_INC)
.cast::<u32>()
.read_unaligned()
};
let info_shift = unsafe {
extensible
.add(layout::INFO_HASH_SHIFT)
.cast::<u32>()
.read_unaligned()
};
let mut hash = hash_uid(uid).wrapping_mul(multiplier);
hash ^= hash >> 33;
let mut want = info_inc.wrapping_add(((hash & INFO_MASK) >> info_shift) as u32);
let mut index = ((hash >> INFO_BITS) as usize) & mask;
for _ in 0..=mask {
let stored = u32::from(unsafe { info.add(index).read() });
if want > stored {
want = want.wrapping_add(info_inc);
index = (index + 1) & mask;
continue;
}
if want < stored {
break;
}
let node = unsafe { key_vals.add(index * NODE_SIZE) };
let key = unsafe { node.cast::<UID>().read_unaligned() };
if key == uid {
let candidate = unsafe { read_usize(node, NODE_VALUE_OFFSET) } as *mut u8;
return if unsafe { extension_id(candidate) } == uid {
candidate
} else {
std::ptr::null_mut()
};
}
want = want.wrapping_add(info_inc);
index = (index + 1) & mask;
}
std::ptr::null_mut()
}
unsafe fn extension_id(candidate: *mut u8) -> UID {
if candidate.is_null() {
return 0;
}
#[cfg(not(target_env = "msvc"))]
type GetIdFn = unsafe extern "C" fn(*mut u8) -> UID;
#[cfg(target_env = "msvc")]
type GetIdFn = unsafe extern "thiscall" fn(*mut u8) -> UID;
let Some((this, f_ptr)) =
(unsafe { super::vtable::secondary_call_target_ptr(candidate, 0, SLOT_GET_EXTENSION_ID) })
else {
return 0;
};
let get_id: GetIdFn = unsafe { std::mem::transmute(f_ptr) };
unsafe { get_id(this) }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hash_matches_the_cpp_implementation() {
assert_eq!(hash_uid(0xbc03_376a_a359_1a11), 0x135c_f77a);
assert_eq!(hash_uid(1), 0x92fd_5b26);
assert_eq!(hash_uid(0xffff_ffff_ffff_ffff), 0x84aa_9ccc);
}
#[test]
fn layout_offsets_match_clangs_record_dump() {
#[cfg(not(target_env = "msvc"))]
assert_eq!(
[
layout::HASH_MULTIPLIER,
layout::KEY_VALS,
layout::INFO,
layout::NUM_ELEMENTS,
layout::MASK,
],
[4, 12, 16, 20, 24]
);
#[cfg(target_env = "msvc")]
assert_eq!(
[
layout::HASH_MULTIPLIER,
layout::KEY_VALS,
layout::INFO,
layout::NUM_ELEMENTS,
layout::MASK,
],
[16, 24, 28, 32, 36]
);
assert_eq!(NODE_SIZE, 16, "UID + IExtension* + bool, padded");
assert_eq!(NODE_VALUE_OFFSET, 8);
}
#[test]
fn a_null_extensible_has_no_extensions() {
assert!(unsafe { extension(std::ptr::null_mut(), 1) }.is_null());
}
}