use crate::hints::unwrap_or_bug_message_hint;
use core::iter::Iterator;
#[cfg(not(feature = "more_numa_nodes"))]
pub const MAX_NUMA_NODES_SUPPORTED_: usize = 64;
#[cfg(feature = "more_numa_nodes")]
pub const MAX_NUMA_NODES_SUPPORTED_: usize = 1024;
pub const MAX_NUMA_NODES_SUPPORTED: usize = MAX_NUMA_NODES_SUPPORTED_;
const NUMA_NODE_TOO_LARGE: &str = "this hardware supports more NUMA-nodes than expected, use the `more_numa_nodes` feature to increase the limit";
pub struct DataPerNUMANodeManager<T>([T; MAX_NUMA_NODES_SUPPORTED]);
impl<T> DataPerNUMANodeManager<T> {
pub const fn from_arr(inner: [T; MAX_NUMA_NODES_SUPPORTED]) -> Self {
Self(inner)
}
pub fn get_ref_by_node(&self, numa_node: usize) -> &T {
unwrap_or_bug_message_hint(self.0.get(numa_node), NUMA_NODE_TOO_LARGE)
}
pub fn get_mut_by_node(&mut self, numa_node: usize) -> &mut T {
unwrap_or_bug_message_hint(self.0.get_mut(numa_node), NUMA_NODE_TOO_LARGE)
}
pub fn iter(&self) -> impl Iterator<Item = &T> {
self.0.iter()
}
pub fn iter_mut(&mut self) -> impl Iterator<Item = &mut T> {
self.0.iter_mut()
}
pub fn as_ptr(&self) -> *const [T; MAX_NUMA_NODES_SUPPORTED] {
self.0.as_ptr().cast()
}
}
impl<T: Default> Default for DataPerNUMANodeManager<T> {
fn default() -> Self {
Self(core::array::from_fn(|_| T::default()))
}
}
pub fn get_current_thread_numa_node() -> usize {
#[cfg(all(target_os = "linux", not(miri)))]
{
use core::mem::MaybeUninit;
let mut numa_node: MaybeUninit<u32> = MaybeUninit::uninit();
unsafe {
libc::syscall(
libc::SYS_getcpu,
core::ptr::null::<libc::c_void>(),
numa_node.as_mut_ptr(),
core::ptr::null::<libc::c_void>(),
);
}
unsafe { numa_node.assume_init() as usize }
}
#[cfg(any(not(target_os = "linux"), miri))]
{
0
}
}
#[cfg(all(test, not(miri)))]
mod tests {
use super::*;
use alloc::vec::Vec;
#[test]
fn test_data_per_numa_node_manager_iterators() {
let mut arr = [1i32; MAX_NUMA_NODES_SUPPORTED];
for (i, item) in arr.iter_mut().enumerate().take(8) {
*item = i32::try_from(i + 1).unwrap();
}
let mut manager = DataPerNUMANodeManager::from_arr(arr);
let values: Vec<i32> = manager.iter().copied().collect();
assert_eq!(values[0], 1);
assert_eq!(values[7], 8);
assert_eq!(values[8], 1);
for val in manager.iter_mut().take(4) {
*val *= 2;
}
assert_eq!(*manager.get_ref_by_node(0), 2);
assert_eq!(*manager.get_ref_by_node(3), 8);
assert_eq!(*manager.get_ref_by_node(4), 5);
let enumerated: Vec<(usize, &i32)> = manager.iter().enumerate().collect();
assert_eq!(enumerated[0], (0, &2));
assert_eq!(enumerated[3], (3, &8));
assert_eq!(enumerated[4], (4, &5));
for (node_id, val) in manager.iter_mut().enumerate() {
if node_id % 2 == 0 {
*val += 10;
}
}
assert_eq!(*manager.get_ref_by_node(0), 12);
assert_eq!(*manager.get_ref_by_node(1), 4);
assert_eq!(*manager.get_ref_by_node(2), 16);
}
#[test]
fn test_get_current_thread_numa_node() {
let node_id = get_current_thread_numa_node();
assert!(node_id < 1024, "node: {node_id}");
}
#[test]
fn test_data_per_numa_node_manager_bounds() {
let manager = DataPerNUMANodeManager::from_arr([0u8; MAX_NUMA_NODES_SUPPORTED]);
for i in 0..MAX_NUMA_NODES_SUPPORTED {
let _ref = manager.get_ref_by_node(i);
}
}
#[test]
fn test_common_case() {
let numa_node = get_current_thread_numa_node();
let manager = DataPerNUMANodeManager::from_arr([0u8; MAX_NUMA_NODES_SUPPORTED]);
assert_eq!(*manager.get_ref_by_node(numa_node), 0);
}
}