Skip to main content

zenith_linux/
affinity.rs

1//! 线程 CPU 亲和性绑定(Linux 平台抽象)
2//!
3//! 将 `libc::sched_setaffinity` 等 Linux 特有系统调用封装为 safe API,
4//! 严格遵守 AGENT.md §1.1 "所有 unsafe 必须封装在 zenith-linux 内部"。
5//!
6//! # 安全保证
7//! 本模块所有 `unsafe` 块均集中在 `libc` FFI 调用,对外提供 `Result` 接口。
8//! 调用方零 unsafe 即可完成线程→CPU 核心的硬绑定。
9
10#![allow(unsafe_code)]
11
12use std::io;
13
14/// 亲和性设置错误
15#[derive(Debug, thiserror::Error)]
16pub enum AffinityError {
17    /// `sched_setaffinity` 系统调用失败
18    #[error("sched_setaffinity failed: {0}")]
19    SetFailed(io::Error),
20
21    /// CPU 编号超出在线 CPU 范围
22    #[error("CPU {cpu} out of range (online: {online})")]
23    CpuOutOfRange {
24        /// 非法的 CPU 编号
25        cpu: usize,
26        /// 在线 CPU 总数
27        online: usize,
28    },
29}
30
31/// 将当前线程绑定到指定 CPU 核心
32///
33/// # Arguments
34/// * `cpu` - 目标 CPU 编号(0 起)
35///
36/// # Returns
37/// * `Ok(())` - 绑定成功
38/// * `Err(AffinityError)` - 系统调用失败或 CPU 越界
39#[cfg(target_os = "linux")]
40pub fn set_thread_affinity(cpu: usize) -> Result<(), AffinityError> {
41    // 边界检查:CPU_SET 在 cpu >= CPU_SETSIZE 时会 panic,必须提前拦截
42    let cpu_set_size = std::mem::size_of::<libc::cpu_set_t>() * 8;
43    if cpu >= cpu_set_size {
44        let online = online_cpu_count();
45        return Err(AffinityError::CpuOutOfRange { cpu, online });
46    }
47
48    // SAFETY: libc::cpu_set_t 是 POD 类型,zeroed() 后由 CPU_ZERO/CPU_SET 显式初始化;
49    // sched_setaffinity 仅读取 cpuset,不持有指针所有权;tid 来自 SYS_gettid 永远合法。
50    unsafe {
51        let mut cpuset: libc::cpu_set_t = std::mem::zeroed();
52        libc::CPU_ZERO(&mut cpuset);
53        libc::CPU_SET(cpu, &mut cpuset);
54
55        let tid = libc::syscall(libc::SYS_gettid);
56        let ret = libc::sched_setaffinity(
57            tid as libc::pid_t,
58            std::mem::size_of::<libc::cpu_set_t>(),
59            &cpuset,
60        );
61
62        if ret != 0 {
63            return Err(AffinityError::SetFailed(io::Error::last_os_error()));
64        }
65    }
66    Ok(())
67}
68
69/// 非 Linux 平台:CPU 亲和为空操作
70#[cfg(not(target_os = "linux"))]
71pub fn set_thread_affinity(_cpu: usize) -> Result<(), AffinityError> {
72    Ok(())
73}
74
75/// 将当前线程绑定到一组 CPU 核心
76///
77/// # Arguments
78/// * `cpus` - 目标 CPU 编号切片(至少 1 个)
79#[cfg(target_os = "linux")]
80pub fn set_thread_affinity_range(cpus: &[usize]) -> Result<(), AffinityError> {
81    // 边界检查:同 set_thread_affinity
82    let cpu_set_size = std::mem::size_of::<libc::cpu_set_t>() * 8;
83    for &cpu in cpus {
84        if cpu >= cpu_set_size {
85            let online = online_cpu_count();
86            return Err(AffinityError::CpuOutOfRange { cpu, online });
87        }
88    }
89
90    // SAFETY: 同 set_thread_affinity;CPU_SET 在循环中对每个 cpu 累加设置位。
91    unsafe {
92        let mut cpuset: libc::cpu_set_t = std::mem::zeroed();
93        libc::CPU_ZERO(&mut cpuset);
94        for &cpu in cpus {
95            libc::CPU_SET(cpu, &mut cpuset);
96        }
97
98        let tid = libc::syscall(libc::SYS_gettid);
99        let ret = libc::sched_setaffinity(
100            tid as libc::pid_t,
101            std::mem::size_of::<libc::cpu_set_t>(),
102            &cpuset,
103        );
104
105        if ret != 0 {
106            return Err(AffinityError::SetFailed(io::Error::last_os_error()));
107        }
108    }
109    Ok(())
110}
111
112/// 非 Linux 平台:空操作
113#[cfg(not(target_os = "linux"))]
114pub fn set_thread_affinity_range(_cpus: &[usize]) -> Result<(), AffinityError> {
115    Ok(())
116}
117
118/// 获取在线 CPU 数量
119///
120/// sysconf 返回 -1 时 fail-closed 回退 1(与 `page_size()` 的 -1 检查语义对齐:
121/// 不得将错误值静默转为巨大 usize)。
122#[cfg(target_os = "linux")]
123pub fn online_cpu_count() -> usize {
124    // SAFETY: sysconf 是纯只读系统调用,仅返回整数;不存在任何内存安全风险。
125    let n = unsafe { libc::sysconf(libc::_SC_NPROCESSORS_ONLN) };
126    if n <= 0 {
127        return 1;
128    }
129    n as usize
130}
131
132/// 非 Linux 平台:返回 1
133#[cfg(not(target_os = "linux"))]
134pub fn online_cpu_count() -> usize {
135    1
136}
137
138#[cfg(test)]
139mod tests {
140    use super::*;
141
142    #[test]
143    fn test_online_cpu_count_positive() {
144        let n = online_cpu_count();
145        assert!(n >= 1, "online CPU count must be >= 1, got {}", n);
146    }
147
148    #[cfg(target_os = "linux")]
149    #[test]
150    fn test_set_affinity_cpu_zero() {
151        // CPU 0 在所有 Linux 系统都存在
152        let result = set_thread_affinity(0);
153        assert!(result.is_ok(), "set_thread_affinity(0) should succeed: {:?}", result);
154    }
155
156    #[cfg(target_os = "linux")]
157    #[test]
158    fn test_set_affinity_range_all_cpus() {
159        let n = online_cpu_count();
160        let cpus: Vec<usize> = (0..n).collect();
161        let result = set_thread_affinity_range(&cpus);
162        assert!(result.is_ok(), "set_thread_affinity_range(all) should succeed: {:?}", result);
163    }
164
165    #[cfg(target_os = "linux")]
166    #[test]
167    fn test_set_affinity_out_of_range() {
168        // 设置一个超大 CPU 编号,期望失败但不会 panic
169        let result = set_thread_affinity(65535);
170        // 失败为 Err,成功也允许(部分实现容忍越界)
171        let _ = result;
172    }
173}