Skip to main content

hermes_simd_core/numa/
binding.rs

1use crate::numa::locality::current_numa_node;
2
3/// RAII scope guard that binds the current thread to a specific NUMA node.
4pub struct NumaBinding {
5    #[cfg(all(target_os = "linux", feature = "libnuma"))]
6    old_mask: *mut core::ffi::c_void,
7    #[cfg(target_os = "windows")]
8    old_mask: usize,
9}
10
11impl NumaBinding {
12    /// Bind the current thread to the specified NUMA node.
13    pub fn bind(node: u32) -> Self {
14        if current_numa_node() == Some(node) {
15            #[cfg(all(target_os = "linux", feature = "libnuma"))]
16            {
17                return Self {
18                    old_mask: core::ptr::null_mut(),
19                };
20            }
21            #[cfg(target_os = "windows")]
22            {
23                return Self { old_mask: 0 };
24            }
25            #[cfg(not(any(all(target_os = "linux", feature = "libnuma"), target_os = "windows")))]
26            {
27                return Self {};
28            }
29        }
30
31        #[cfg(all(target_os = "linux", feature = "libnuma"))]
32        {
33            #[link(name = "numa")]
34            extern "C" {
35                fn numa_allocate_nodemask() -> *mut core::ffi::c_void;
36                fn numa_bitmask_setbit(mask: *mut core::ffi::c_void, bit: u32);
37                fn numa_bind(mask: *mut core::ffi::c_void);
38                fn numa_get_run_node_mask() -> *mut core::ffi::c_void;
39                fn numa_bitmask_free(mask: *mut core::ffi::c_void);
40            }
41            unsafe {
42                let old = numa_get_run_node_mask();
43                let mask = numa_allocate_nodemask();
44                if !mask.is_null() {
45                    numa_bitmask_setbit(mask, node);
46                    numa_bind(mask);
47                    numa_bitmask_free(mask);
48                }
49                Self { old_mask: old }
50            }
51        }
52        #[cfg(target_os = "windows")]
53        {
54            extern "system" {
55                fn GetCurrentThread() -> *mut core::ffi::c_void;
56                fn SetThreadAffinityMask(
57                    hThread: *mut core::ffi::c_void,
58                    dwThreadAffinityMask: usize,
59                ) -> usize;
60                fn GetNumaNodeProcessorMask(node: u8, processor_mask: *mut u64) -> i32;
61            }
62            unsafe {
63                let thread = GetCurrentThread();
64                let mut mask = 0u64;
65                if GetNumaNodeProcessorMask(node as u8, &mut mask) != 0 && mask != 0 {
66                    let old = SetThreadAffinityMask(thread, mask as usize);
67                    Self { old_mask: old }
68                } else {
69                    Self { old_mask: 0 }
70                }
71            }
72        }
73        #[cfg(not(any(all(target_os = "linux", feature = "libnuma"), target_os = "windows")))]
74        {
75            let _ = node;
76            Self {}
77        }
78    }
79}
80
81#[cfg(any(all(target_os = "linux", feature = "libnuma"), target_os = "windows"))]
82impl Drop for NumaBinding {
83    fn drop(&mut self) {
84        #[cfg(all(target_os = "linux", feature = "libnuma"))]
85        {
86            if !self.old_mask.is_null() {
87                #[link(name = "numa")]
88                extern "C" {
89                    fn numa_bind(mask: *mut core::ffi::c_void);
90                    fn numa_bitmask_free(mask: *mut core::ffi::c_void);
91                }
92                unsafe {
93                    numa_bind(self.old_mask);
94                    numa_bitmask_free(self.old_mask);
95                }
96            }
97        }
98        #[cfg(target_os = "windows")]
99        {
100            if self.old_mask != 0 {
101                extern "system" {
102                    fn GetCurrentThread() -> *mut core::ffi::c_void;
103                    fn SetThreadAffinityMask(
104                        hThread: *mut core::ffi::c_void,
105                        dwThreadAffinityMask: usize,
106                    ) -> usize;
107                }
108                unsafe {
109                    SetThreadAffinityMask(GetCurrentThread(), self.old_mask);
110                }
111            }
112        }
113    }
114}