Skip to main content

tenferro_cpu/
affinity.rs

1use std::num::NonZeroUsize;
2
3use thiserror::Error;
4
5use crate::{CpuId, CpuSet};
6
7/// Typed failures from the operating-system CPU-affinity boundary.
8///
9/// These errors stay CPU-local until a context-construction failure is
10/// reported to the tensor API, where the complete value is retained as the
11/// source of [`crate::CpuContextError`].
12///
13/// # Examples
14///
15/// ```
16/// use tenferro_cpu::CpuAffinityError;
17///
18/// let error = CpuAffinityError::UnsupportedPlatform;
19/// assert!(error.to_string().contains("unsupported"));
20/// ```
21#[derive(Debug, Error)]
22pub enum CpuAffinityError {
23    /// Constructing the one-CPU mask failed because the CPU set was invalid.
24    #[error("CPU affinity set is invalid: {0}")]
25    CpuSet(#[from] crate::CpuSetError),
26    /// The operating system rejected the requested affinity mask.
27    #[error("setting thread affinity failed: {source}")]
28    Set {
29        #[source]
30        source: std::io::Error,
31    },
32    /// Querying the current worker affinity failed.
33    #[error("querying worker affinity failed: {source}")]
34    Query {
35        #[source]
36        source: std::io::Error,
37    },
38    /// The platform has no supported thread-affinity implementation.
39    #[error("setting thread affinity is unsupported on this platform")]
40    UnsupportedPlatform,
41    /// The operating system did not expose a process affinity mask.
42    #[error("failed to verify worker affinity")]
43    VerificationUnavailable,
44    /// The returned affinity did not contain exactly the requested worker.
45    #[error("verification returned affinity {observed:?}")]
46    Verification { observed: Vec<CpuId> },
47    /// The requested CPU would overflow the affinity-mask size calculation.
48    #[error("affinity mask size overflow")]
49    MaskSizeOverflow,
50    /// The requested CPU exceeds the supported mask allocation limit.
51    #[error("CPU {cpu} exceeds supported affinity mask size of {max_bytes} bytes")]
52    MaskTooLarge { cpu: CpuId, max_bytes: usize },
53    /// Allocating the operating-system affinity mask failed.
54    #[error("failed to allocate affinity mask: {source}")]
55    MaskAllocation {
56        #[source]
57        source: std::collections::TryReserveError,
58    },
59    /// The affinity mask had no CPU entries.
60    #[error("cannot set an empty affinity mask")]
61    EmptyMask,
62}
63
64pub(crate) trait ThreadAffinity: Clone + Send + Sync + 'static {
65    /// Confine the calling thread to `cpus` and report the resulting mask.
66    fn confine_current(&self, cpus: &CpuSet) -> Result<CpuSet, CpuAffinityError>;
67}
68
69#[derive(Clone, Copy, Debug)]
70pub(crate) struct SystemThreadAffinity;
71
72impl ThreadAffinity for SystemThreadAffinity {
73    fn confine_current(&self, cpus: &CpuSet) -> Result<CpuSet, CpuAffinityError> {
74        set_current_thread_affinity(cpus)?;
75        process_cpu_affinity().ok_or(CpuAffinityError::VerificationUnavailable)
76    }
77}
78
79#[cfg(all(test, target_os = "linux"))]
80pub(crate) fn current_cpu() -> Result<CpuId, CpuAffinityError> {
81    unsafe extern "C" {
82        fn sched_getcpu() -> i32;
83    }
84    // SAFETY: `sched_getcpu` takes no arguments and returns the calling
85    // thread's current logical CPU or a negative error sentinel.
86    let cpu = unsafe { sched_getcpu() };
87    usize::try_from(cpu)
88        .map(CpuId::new)
89        .map_err(|_| CpuAffinityError::Query {
90            source: std::io::Error::last_os_error(),
91        })
92}
93
94fn set_current_thread_affinity(cpus: &CpuSet) -> Result<(), CpuAffinityError> {
95    #[cfg(any(target_os = "linux", target_os = "android"))]
96    {
97        unsafe extern "C" {
98            fn sched_setaffinity(
99                pid: i32,
100                cpusetsize: usize,
101                mask: *const core::ffi::c_void,
102            ) -> i32;
103        }
104
105        let mask = build_affinity_mask(cpus)?;
106        // SAFETY: `mask` remains allocated for the call, `cpusetsize` exactly
107        // matches its byte length, and pid 0 selects the calling thread.
108        let rc =
109            unsafe { sched_setaffinity(0, mask.len(), mask.as_ptr().cast::<core::ffi::c_void>()) };
110        (rc == 0)
111            .then_some(())
112            .ok_or_else(|| CpuAffinityError::Set {
113                source: std::io::Error::last_os_error(),
114            })
115    }
116    #[cfg(not(any(target_os = "linux", target_os = "android")))]
117    {
118        let _ = cpus;
119        Err(CpuAffinityError::UnsupportedPlatform)
120    }
121}
122
123#[cfg(any(target_os = "linux", target_os = "android", test))]
124fn build_affinity_mask(cpus: &CpuSet) -> Result<Vec<u8>, CpuAffinityError> {
125    const MIN_MASK_BYTES: usize = 128;
126    const MAX_MASK_BYTES: usize = 1 << 20;
127
128    let highest_cpu = cpus
129        .as_slice()
130        .last()
131        .copied()
132        .ok_or(CpuAffinityError::EmptyMask)?;
133    let required_bytes = highest_cpu
134        .as_usize()
135        .checked_div(u8::BITS as usize)
136        .and_then(|index| index.checked_add(1))
137        .ok_or(CpuAffinityError::MaskSizeOverflow)?
138        .max(MIN_MASK_BYTES);
139    if required_bytes > MAX_MASK_BYTES {
140        return Err(CpuAffinityError::MaskTooLarge {
141            cpu: highest_cpu,
142            max_bytes: MAX_MASK_BYTES,
143        });
144    }
145
146    let mut mask = Vec::new();
147    mask.try_reserve_exact(required_bytes)
148        .map_err(|source| CpuAffinityError::MaskAllocation { source })?;
149    mask.resize(required_bytes, 0u8);
150    for cpu in cpus.as_slice() {
151        let byte = cpu.as_usize() / u8::BITS as usize;
152        let bit = cpu.as_usize() % u8::BITS as usize;
153        mask[byte] |= 1 << bit;
154    }
155    Ok(mask)
156}
157
158/// Return a best-effort CPU count available to the current process.
159///
160/// This first tries an OS-standard process-affinity query when supported, then
161/// falls back to `std::thread::available_parallelism()`, and finally to `1`.
162///
163/// # Examples
164///
165/// ```
166/// let available = tenferro_cpu::available_parallelism();
167/// assert!(available >= 1);
168/// ```
169pub fn available_parallelism() -> usize {
170    process_cpu_affinity_count()
171        .or_else(standard_available_parallelism)
172        .unwrap_or(1)
173}
174
175/// Return the current process affinity mask size when the platform exposes a
176/// standard affinity API.
177///
178/// Platforms without an affinity query return `None`.
179///
180/// # Examples
181///
182/// ```
183/// let count = tenferro_cpu::process_cpu_affinity_count();
184/// if let Some(count) = count {
185///     assert!(count >= 1);
186/// }
187/// ```
188pub fn process_cpu_affinity_count() -> Option<usize> {
189    platform_process_cpu_affinity_count()
190}
191
192/// Return the process affinity mask as logical CPU identifiers when supported.
193///
194/// The returned set preserves sparse operating-system CPU IDs. Platforms where
195/// the standard affinity API exposes only a count return `None`.
196///
197/// # Examples
198///
199/// ```
200/// if let Some(cpus) = tenferro_cpu::process_cpu_affinity() {
201///     assert!(!cpus.is_empty());
202/// }
203/// ```
204pub fn process_cpu_affinity() -> Option<CpuSet> {
205    platform_process_cpu_affinity()
206}
207
208pub(crate) fn standard_available_parallelism() -> Option<usize> {
209    std::thread::available_parallelism()
210        .ok()
211        .map(NonZeroUsize::get)
212}
213
214#[cfg(test)]
215fn count_affinity_mask_bits(mask: &[u8]) -> Option<usize> {
216    cpu_set_from_affinity_mask(mask).map(|cpus| cpus.len())
217}
218
219#[cfg(any(target_os = "linux", target_os = "android", test))]
220fn cpu_set_from_affinity_mask(mask: &[u8]) -> Option<CpuSet> {
221    let cpus = mask.iter().enumerate().flat_map(|(byte_index, byte)| {
222        (0..u8::BITS as usize)
223            .filter(move |bit| byte & (1 << bit) != 0)
224            .map(move |bit| CpuId::new(byte_index * u8::BITS as usize + bit))
225    });
226    CpuSet::new(cpus).ok()
227}
228
229#[cfg(any(target_os = "linux", target_os = "android"))]
230const LINUX_EINVAL: i32 = 22;
231
232#[cfg(any(target_os = "linux", target_os = "android"))]
233fn linux_next_affinity_mask_bytes(mask_bytes: usize, errno: Option<i32>) -> Option<usize> {
234    (errno == Some(LINUX_EINVAL))
235        .then(|| mask_bytes.checked_mul(2))
236        .flatten()
237}
238
239#[cfg(any(target_os = "linux", target_os = "android"))]
240fn platform_process_cpu_affinity_count() -> Option<usize> {
241    platform_process_cpu_affinity().map(|cpus| cpus.len())
242}
243
244#[cfg(any(target_os = "linux", target_os = "android"))]
245fn platform_process_cpu_affinity() -> Option<CpuSet> {
246    unsafe extern "C" {
247        fn sched_getaffinity(pid: i32, cpusetsize: usize, mask: *mut core::ffi::c_void) -> i32;
248    }
249
250    const INITIAL_MASK_BYTES: usize = 128;
251
252    let mut mask_bytes = INITIAL_MASK_BYTES;
253    loop {
254        let mut mask = vec![0u8; mask_bytes];
255        // SAFETY: `mask` is a live allocation of `mask_bytes` bytes, and pid 0
256        // asks the OS to query the current process affinity.
257        let rc = unsafe {
258            sched_getaffinity(0, mask_bytes, mask.as_mut_ptr().cast::<core::ffi::c_void>())
259        };
260        if rc == 0 {
261            return cpu_set_from_affinity_mask(&mask);
262        }
263
264        mask_bytes = linux_next_affinity_mask_bytes(
265            mask_bytes,
266            std::io::Error::last_os_error().raw_os_error(),
267        )?;
268    }
269}
270
271#[cfg(not(any(target_os = "linux", target_os = "android")))]
272fn platform_process_cpu_affinity() -> Option<CpuSet> {
273    None
274}
275
276#[cfg(target_os = "windows")]
277fn platform_process_cpu_affinity_count() -> Option<usize> {
278    type Handle = *mut core::ffi::c_void;
279    type DwordPtr = usize;
280    type Word = u16;
281
282    unsafe extern "system" {
283        fn GetCurrentProcess() -> Handle;
284        fn GetProcessAffinityMask(
285            process: Handle,
286            process_affinity_mask: *mut DwordPtr,
287            system_affinity_mask: *mut DwordPtr,
288        ) -> i32;
289        fn GetActiveProcessorGroupCount() -> Word;
290        fn GetActiveProcessorCount(group_number: Word) -> u32;
291        fn GetProcessGroupAffinity(
292            process: Handle,
293            group_count: *mut Word,
294            group_array: *mut Word,
295        ) -> i32;
296    }
297
298    // SAFETY: `GetCurrentProcess` takes no arguments and returns a pseudo-handle
299    // owned by the process; it must not be closed by the caller.
300    let process = unsafe { GetCurrentProcess() };
301    // SAFETY: This Windows query takes no pointers and has no preconditions.
302    let system_group_count = unsafe { GetActiveProcessorGroupCount() };
303
304    if system_group_count <= 1 {
305        let mut process_mask = 0usize;
306        let mut system_mask = 0usize;
307        // SAFETY: `process` is the current-process pseudo-handle and both
308        // output pointers refer to live local variables for the duration of the call.
309        let ok = unsafe {
310            GetProcessAffinityMask(
311                process,
312                std::ptr::addr_of_mut!(process_mask),
313                std::ptr::addr_of_mut!(system_mask),
314            )
315        };
316        if ok != 0 {
317            let count = process_mask.count_ones() as usize;
318            return (count > 0).then_some(count);
319        }
320        // SAFETY: Group 0 exists when Windows reports at most one active group.
321        let count = unsafe { GetActiveProcessorCount(0) } as usize;
322        return (count > 0).then_some(count);
323    }
324
325    let mut group_count: Word = 0;
326    // SAFETY: Windows accepts a null group array to query the required processor-group
327    // count. That probe is the expected failure path: `group_count` is a live output
328    // variable and a nonzero count after a failed call is the value needed for the
329    // second call below.
330    let ok = unsafe {
331        GetProcessGroupAffinity(
332            process,
333            std::ptr::addr_of_mut!(group_count),
334            std::ptr::null_mut(),
335        )
336    };
337    if ok != 0 || group_count == 0 {
338        // SAFETY: `u16::MAX` requests the total count across all processor groups.
339        let count = unsafe { GetActiveProcessorCount(u16::MAX) } as usize;
340        return (count > 0).then_some(count);
341    }
342
343    let mut groups = vec![0u16; group_count as usize];
344    // SAFETY: `groups` has `group_count` entries and both output pointers stay
345    // valid for the duration of the call.
346    let ok = unsafe {
347        GetProcessGroupAffinity(
348            process,
349            std::ptr::addr_of_mut!(group_count),
350            groups.as_mut_ptr(),
351        )
352    };
353    if ok == 0 || group_count == 0 {
354        // SAFETY: `u16::MAX` requests the total count across all processor groups.
355        let count = unsafe { GetActiveProcessorCount(u16::MAX) } as usize;
356        return (count > 0).then_some(count);
357    }
358
359    if group_count == 1 {
360        let mut process_mask = 0usize;
361        let mut system_mask = 0usize;
362        // SAFETY: `process` is the current-process pseudo-handle and both
363        // output pointers refer to live local variables for the duration of the call.
364        let ok = unsafe {
365            GetProcessAffinityMask(
366                process,
367                std::ptr::addr_of_mut!(process_mask),
368                std::ptr::addr_of_mut!(system_mask),
369            )
370        };
371        if ok != 0 {
372            let count = process_mask.count_ones() as usize;
373            return (count > 0).then_some(count);
374        }
375    }
376
377    let count = groups
378        .into_iter()
379        .map(|group| {
380            // SAFETY: Group identifiers are returned by `GetProcessGroupAffinity`.
381            (unsafe { GetActiveProcessorCount(group) }) as usize
382        })
383        .sum();
384    (count > 0).then_some(count)
385}
386
387#[cfg(not(any(target_os = "linux", target_os = "android", target_os = "windows")))]
388fn platform_process_cpu_affinity_count() -> Option<usize> {
389    None
390}
391
392#[cfg(test)]
393mod tests;