1use crate::hints::unwrap_or_bug_message_hint;
21use core::iter::Iterator;
22
23#[cfg(not(feature = "more_numa_nodes"))]
24pub const MAX_NUMA_NODES_SUPPORTED_: usize = 64;
25#[cfg(feature = "more_numa_nodes")]
26pub const MAX_NUMA_NODES_SUPPORTED_: usize = 1024;
27
28pub const MAX_NUMA_NODES_SUPPORTED: usize = MAX_NUMA_NODES_SUPPORTED_;
33
34const NUMA_NODE_TOO_LARGE: &str = "this hardware supports more NUMA-nodes than expected, use the `more_numa_nodes` feature to increase the limit";
35
36pub struct DataPerNUMANodeManager<T>([T; MAX_NUMA_NODES_SUPPORTED]);
51
52impl<T> DataPerNUMANodeManager<T> {
53 pub const fn from_arr(inner: [T; MAX_NUMA_NODES_SUPPORTED]) -> Self {
55 Self(inner)
56 }
57
58 pub fn get_ref_by_node(&self, numa_node: usize) -> &T {
64 unwrap_or_bug_message_hint(self.0.get(numa_node), NUMA_NODE_TOO_LARGE)
65 }
66
67 pub fn get_mut_by_node(&mut self, numa_node: usize) -> &mut T {
73 unwrap_or_bug_message_hint(self.0.get_mut(numa_node), NUMA_NODE_TOO_LARGE)
74 }
75
76 pub fn iter(&self) -> impl Iterator<Item = &T> {
78 self.0.iter()
79 }
80
81 pub fn iter_mut(&mut self) -> impl Iterator<Item = &mut T> {
83 self.0.iter_mut()
84 }
85
86 pub fn as_ptr(&self) -> *const [T; MAX_NUMA_NODES_SUPPORTED] {
88 self.0.as_ptr().cast()
89 }
90}
91
92impl<T: Default> Default for DataPerNUMANodeManager<T> {
93 fn default() -> Self {
94 Self(core::array::from_fn(|_| T::default()))
95 }
96}
97
98pub fn get_current_thread_numa_node() -> usize {
112 #[cfg(all(target_os = "linux", not(miri)))]
113 {
114 use core::mem::MaybeUninit;
115
116 let mut numa_node: MaybeUninit<u32> = MaybeUninit::uninit();
117
118 unsafe {
119 libc::syscall(
120 libc::SYS_getcpu,
121 core::ptr::null::<libc::c_void>(),
122 numa_node.as_mut_ptr(),
123 core::ptr::null::<libc::c_void>(),
124 );
125 }
126
127 unsafe { numa_node.assume_init() as usize }
128 }
129
130 #[cfg(any(not(target_os = "linux"), miri))]
131 {
132 0
133 }
134}
135
136#[cfg(all(test, not(miri)))]
137mod tests {
138 use super::*;
139 use alloc::vec::Vec;
140
141 #[test]
142 fn test_data_per_numa_node_manager_iterators() {
143 let mut arr = [1i32; MAX_NUMA_NODES_SUPPORTED];
144 for (i, item) in arr.iter_mut().enumerate().take(8) {
145 *item = i32::try_from(i + 1).unwrap();
146 }
147 let mut manager = DataPerNUMANodeManager::from_arr(arr);
148
149 let values: Vec<i32> = manager.iter().copied().collect();
151 assert_eq!(values[0], 1);
152 assert_eq!(values[7], 8);
153 assert_eq!(values[8], 1); for val in manager.iter_mut().take(4) {
157 *val *= 2;
158 }
159 assert_eq!(*manager.get_ref_by_node(0), 2);
160 assert_eq!(*manager.get_ref_by_node(3), 8);
161 assert_eq!(*manager.get_ref_by_node(4), 5);
162
163 let enumerated: Vec<(usize, &i32)> = manager.iter().enumerate().collect();
165 assert_eq!(enumerated[0], (0, &2));
166 assert_eq!(enumerated[3], (3, &8));
167 assert_eq!(enumerated[4], (4, &5));
168
169 for (node_id, val) in manager.iter_mut().enumerate() {
171 if node_id % 2 == 0 {
172 *val += 10;
173 }
174 }
175 assert_eq!(*manager.get_ref_by_node(0), 12);
176 assert_eq!(*manager.get_ref_by_node(1), 4);
177 assert_eq!(*manager.get_ref_by_node(2), 16);
178 }
179
180 #[test]
181 fn test_get_current_thread_numa_node() {
182 let node_id = get_current_thread_numa_node();
183 assert!(node_id < 1024, "node: {node_id}");
184 }
185
186 #[test]
187 fn test_data_per_numa_node_manager_bounds() {
188 let manager = DataPerNUMANodeManager::from_arr([0u8; MAX_NUMA_NODES_SUPPORTED]);
189
190 for i in 0..MAX_NUMA_NODES_SUPPORTED {
192 let _ref = manager.get_ref_by_node(i);
193 }
194 }
195
196 #[test]
197 fn test_common_case() {
198 let numa_node = get_current_thread_numa_node();
199 let manager = DataPerNUMANodeManager::from_arr([0u8; MAX_NUMA_NODES_SUPPORTED]);
200
201 assert_eq!(*manager.get_ref_by_node(numa_node), 0);
202 }
203}