use std::num::NonZeroUsize;
use thiserror::Error;
use crate::{CpuId, CpuSet};
#[derive(Debug, Error)]
pub enum CpuAffinityError {
#[error("CPU affinity set is invalid: {0}")]
CpuSet(#[from] crate::CpuSetError),
#[error("setting thread affinity failed: {source}")]
Set {
#[source]
source: std::io::Error,
},
#[error("querying worker affinity failed: {source}")]
Query {
#[source]
source: std::io::Error,
},
#[error("setting thread affinity is unsupported on this platform")]
UnsupportedPlatform,
#[error("failed to verify worker affinity")]
VerificationUnavailable,
#[error("verification returned affinity {observed:?}")]
Verification { observed: Vec<CpuId> },
#[error("affinity mask size overflow")]
MaskSizeOverflow,
#[error("CPU {cpu} exceeds supported affinity mask size of {max_bytes} bytes")]
MaskTooLarge { cpu: CpuId, max_bytes: usize },
#[error("failed to allocate affinity mask: {source}")]
MaskAllocation {
#[source]
source: std::collections::TryReserveError,
},
#[error("cannot set an empty affinity mask")]
EmptyMask,
}
pub(crate) trait ThreadAffinity: Clone + Send + Sync + 'static {
fn pin_current(&self, cpu: CpuId) -> Result<CpuSet, CpuAffinityError>;
}
#[derive(Clone, Copy, Debug)]
pub(crate) struct SystemThreadAffinity;
impl ThreadAffinity for SystemThreadAffinity {
fn pin_current(&self, cpu: CpuId) -> Result<CpuSet, CpuAffinityError> {
set_current_thread_affinity(&CpuSet::new([cpu])?)?;
process_cpu_affinity().ok_or(CpuAffinityError::VerificationUnavailable)
}
}
#[cfg(all(test, target_os = "linux"))]
pub(crate) fn current_cpu() -> Result<CpuId, CpuAffinityError> {
unsafe extern "C" {
fn sched_getcpu() -> i32;
}
let cpu = unsafe { sched_getcpu() };
usize::try_from(cpu)
.map(CpuId::new)
.map_err(|_| CpuAffinityError::Query {
source: std::io::Error::last_os_error(),
})
}
fn set_current_thread_affinity(cpus: &CpuSet) -> Result<(), CpuAffinityError> {
#[cfg(any(target_os = "linux", target_os = "android"))]
{
unsafe extern "C" {
fn sched_setaffinity(
pid: i32,
cpusetsize: usize,
mask: *const core::ffi::c_void,
) -> i32;
}
let mask = build_affinity_mask(cpus)?;
let rc =
unsafe { sched_setaffinity(0, mask.len(), mask.as_ptr().cast::<core::ffi::c_void>()) };
(rc == 0)
.then_some(())
.ok_or_else(|| CpuAffinityError::Set {
source: std::io::Error::last_os_error(),
})
}
#[cfg(not(any(target_os = "linux", target_os = "android")))]
{
let _ = cpus;
Err(CpuAffinityError::UnsupportedPlatform)
}
}
#[cfg(any(target_os = "linux", target_os = "android", test))]
fn build_affinity_mask(cpus: &CpuSet) -> Result<Vec<u8>, CpuAffinityError> {
const MIN_MASK_BYTES: usize = 128;
const MAX_MASK_BYTES: usize = 1 << 20;
let highest_cpu = cpus
.as_slice()
.last()
.copied()
.ok_or(CpuAffinityError::EmptyMask)?;
let required_bytes = highest_cpu
.as_usize()
.checked_div(u8::BITS as usize)
.and_then(|index| index.checked_add(1))
.ok_or(CpuAffinityError::MaskSizeOverflow)?
.max(MIN_MASK_BYTES);
if required_bytes > MAX_MASK_BYTES {
return Err(CpuAffinityError::MaskTooLarge {
cpu: highest_cpu,
max_bytes: MAX_MASK_BYTES,
});
}
let mut mask = Vec::new();
mask.try_reserve_exact(required_bytes)
.map_err(|source| CpuAffinityError::MaskAllocation { source })?;
mask.resize(required_bytes, 0u8);
for cpu in cpus.as_slice() {
let byte = cpu.as_usize() / u8::BITS as usize;
let bit = cpu.as_usize() % u8::BITS as usize;
mask[byte] |= 1 << bit;
}
Ok(mask)
}
pub fn available_parallelism() -> usize {
process_cpu_affinity_count()
.or_else(standard_available_parallelism)
.unwrap_or(1)
}
pub fn process_cpu_affinity_count() -> Option<usize> {
platform_process_cpu_affinity_count()
}
pub fn process_cpu_affinity() -> Option<CpuSet> {
platform_process_cpu_affinity()
}
pub(crate) fn standard_available_parallelism() -> Option<usize> {
std::thread::available_parallelism()
.ok()
.map(NonZeroUsize::get)
}
#[cfg(test)]
fn count_affinity_mask_bits(mask: &[u8]) -> Option<usize> {
cpu_set_from_affinity_mask(mask).map(|cpus| cpus.len())
}
#[cfg(any(target_os = "linux", target_os = "android", test))]
fn cpu_set_from_affinity_mask(mask: &[u8]) -> Option<CpuSet> {
let cpus = mask.iter().enumerate().flat_map(|(byte_index, byte)| {
(0..u8::BITS as usize)
.filter(move |bit| byte & (1 << bit) != 0)
.map(move |bit| CpuId::new(byte_index * u8::BITS as usize + bit))
});
CpuSet::new(cpus).ok()
}
#[cfg(any(target_os = "linux", target_os = "android"))]
const LINUX_EINVAL: i32 = 22;
#[cfg(any(target_os = "linux", target_os = "android"))]
fn linux_next_affinity_mask_bytes(mask_bytes: usize, errno: Option<i32>) -> Option<usize> {
(errno == Some(LINUX_EINVAL))
.then(|| mask_bytes.checked_mul(2))
.flatten()
}
#[cfg(any(target_os = "linux", target_os = "android"))]
fn platform_process_cpu_affinity_count() -> Option<usize> {
platform_process_cpu_affinity().map(|cpus| cpus.len())
}
#[cfg(any(target_os = "linux", target_os = "android"))]
fn platform_process_cpu_affinity() -> Option<CpuSet> {
unsafe extern "C" {
fn sched_getaffinity(pid: i32, cpusetsize: usize, mask: *mut core::ffi::c_void) -> i32;
}
const INITIAL_MASK_BYTES: usize = 128;
let mut mask_bytes = INITIAL_MASK_BYTES;
loop {
let mut mask = vec![0u8; mask_bytes];
let rc = unsafe {
sched_getaffinity(0, mask_bytes, mask.as_mut_ptr().cast::<core::ffi::c_void>())
};
if rc == 0 {
return cpu_set_from_affinity_mask(&mask);
}
mask_bytes = linux_next_affinity_mask_bytes(
mask_bytes,
std::io::Error::last_os_error().raw_os_error(),
)?;
}
}
#[cfg(not(any(target_os = "linux", target_os = "android")))]
fn platform_process_cpu_affinity() -> Option<CpuSet> {
None
}
#[cfg(target_os = "windows")]
fn platform_process_cpu_affinity_count() -> Option<usize> {
type Handle = *mut core::ffi::c_void;
type DwordPtr = usize;
type Word = u16;
unsafe extern "system" {
fn GetCurrentProcess() -> Handle;
fn GetProcessAffinityMask(
process: Handle,
process_affinity_mask: *mut DwordPtr,
system_affinity_mask: *mut DwordPtr,
) -> i32;
fn GetActiveProcessorGroupCount() -> Word;
fn GetActiveProcessorCount(group_number: Word) -> u32;
fn GetProcessGroupAffinity(
process: Handle,
group_count: *mut Word,
group_array: *mut Word,
) -> i32;
}
let process = unsafe { GetCurrentProcess() };
let system_group_count = unsafe { GetActiveProcessorGroupCount() };
if system_group_count <= 1 {
let mut process_mask = 0usize;
let mut system_mask = 0usize;
let ok = unsafe {
GetProcessAffinityMask(
process,
std::ptr::addr_of_mut!(process_mask),
std::ptr::addr_of_mut!(system_mask),
)
};
if ok != 0 {
let count = process_mask.count_ones() as usize;
return (count > 0).then_some(count);
}
let count = unsafe { GetActiveProcessorCount(0) } as usize;
return (count > 0).then_some(count);
}
let mut group_count: Word = 0;
let ok = unsafe {
GetProcessGroupAffinity(
process,
std::ptr::addr_of_mut!(group_count),
std::ptr::null_mut(),
)
};
if ok != 0 || group_count == 0 {
let count = unsafe { GetActiveProcessorCount(u16::MAX) } as usize;
return (count > 0).then_some(count);
}
let mut groups = vec![0u16; group_count as usize];
let ok = unsafe {
GetProcessGroupAffinity(
process,
std::ptr::addr_of_mut!(group_count),
groups.as_mut_ptr(),
)
};
if ok == 0 || group_count == 0 {
let count = unsafe { GetActiveProcessorCount(u16::MAX) } as usize;
return (count > 0).then_some(count);
}
if group_count == 1 {
let mut process_mask = 0usize;
let mut system_mask = 0usize;
let ok = unsafe {
GetProcessAffinityMask(
process,
std::ptr::addr_of_mut!(process_mask),
std::ptr::addr_of_mut!(system_mask),
)
};
if ok != 0 {
let count = process_mask.count_ones() as usize;
return (count > 0).then_some(count);
}
}
let count = groups
.into_iter()
.map(|group| {
(unsafe { GetActiveProcessorCount(group) }) as usize
})
.sum();
(count > 0).then_some(count)
}
#[cfg(not(any(target_os = "linux", target_os = "android", target_os = "windows")))]
fn platform_process_cpu_affinity_count() -> Option<usize> {
None
}
#[cfg(test)]
mod tests;