use std::mem::MaybeUninit;
use crate::bpf_intf;
use crate::bpf_intf::*;
use crate::bpf_skel::*;
use std::ffi::c_int;
use std::ffi::c_ulong;
use std::ffi::CStr;
use std::sync::atomic::AtomicBool;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use std::sync::Once;
use anyhow::bail;
use anyhow::Context;
use anyhow::Result;
use plain::Plain;
use procfs::process::all_processes;
use libbpf_rs::libbpf_sys::bpf_object_open_opts;
use libbpf_rs::OpenObject;
use libbpf_rs::ProgramInput;
use libc::{c_char, pthread_self, pthread_setschedparam, sched_param};
#[cfg(target_env = "musl")]
use libc::timespec;
use scx_utils::compat;
use scx_utils::scx_ops_attach;
use scx_utils::scx_ops_load;
use scx_utils::scx_ops_open;
use scx_utils::uei_exited;
use scx_utils::uei_report;
use scx_utils::Topology;
use scx_utils::UserExitInfo;
use scx_rustland_core::ALLOCATOR;
const SCHED_EXT: i32 = 7;
const TASK_COMM_LEN: usize = 16;
#[allow(dead_code)]
pub const RL_CPU_ANY: i32 = bpf_intf::RL_CPU_ANY as i32;
#[derive(Debug, PartialEq, Eq, PartialOrd, Clone)]
pub struct QueuedTask {
pub pid: i32, pub cpu: i32, pub nr_cpus_allowed: u64, pub flags: u64, pub start_ts: u64, pub stop_ts: u64, pub exec_runtime: u64, pub weight: u64, pub vtime: u64, pub enq_cnt: u64,
pub comm: [c_char; TASK_COMM_LEN], }
impl QueuedTask {
#[allow(dead_code)]
pub fn comm_str(&self) -> String {
let c_str = unsafe { CStr::from_ptr(self.comm.as_ptr()) };
c_str.to_string_lossy().into_owned()
}
}
#[derive(Debug, PartialEq, Eq, PartialOrd, Clone)]
pub struct DispatchedTask {
pub pid: i32, pub cpu: i32, pub flags: u64, pub slice_ns: u64, pub vtime: u64, pub enq_cnt: u64,
}
impl DispatchedTask {
pub fn new(task: &QueuedTask) -> Self {
DispatchedTask {
pid: task.pid,
cpu: task.cpu,
flags: task.flags,
slice_ns: 0, vtime: 0,
enq_cnt: task.enq_cnt,
}
}
}
unsafe impl Plain for bpf_intf::dispatched_task_ctx {}
impl AsMut<bpf_intf::dispatched_task_ctx> for bpf_intf::dispatched_task_ctx {
fn as_mut(&mut self) -> &mut bpf_intf::dispatched_task_ctx {
self
}
}
struct EnqueuedMessage {
inner: bpf_intf::queued_task_ctx,
}
impl EnqueuedMessage {
fn from_bytes(bytes: &[u8]) -> Self {
let queued_task_struct = unsafe { *(bytes.as_ptr() as *const bpf_intf::queued_task_ctx) };
EnqueuedMessage {
inner: queued_task_struct,
}
}
fn to_queued_task(&self) -> QueuedTask {
QueuedTask {
pid: self.inner.pid,
cpu: self.inner.cpu,
nr_cpus_allowed: self.inner.nr_cpus_allowed,
flags: self.inner.flags,
start_ts: self.inner.start_ts,
stop_ts: self.inner.stop_ts,
exec_runtime: self.inner.exec_runtime,
weight: self.inner.weight,
vtime: self.inner.vtime,
enq_cnt: self.inner.enq_cnt,
comm: self.inner.comm,
}
}
}
pub struct BpfScheduler<'cb> {
pub skel: BpfSkel<'cb>, shutdown: Arc<AtomicBool>, queued: libbpf_rs::RingBuffer<'cb>, dispatched: libbpf_rs::UserRingBuffer, struct_ops: Option<libbpf_rs::Link>, }
const BUFSIZE: usize = size_of::<queued_task_ctx>();
#[repr(align(8))]
struct AlignedBuffer([u8; BUFSIZE]);
static mut BUF: AlignedBuffer = AlignedBuffer([0; BUFSIZE]);
static SET_HANDLER: Once = Once::new();
fn set_ctrlc_handler(shutdown: Arc<AtomicBool>) -> Result<(), anyhow::Error> {
SET_HANDLER.call_once(|| {
let shutdown_clone = shutdown.clone();
ctrlc::set_handler(move || {
shutdown_clone.store(true, Ordering::Relaxed);
})
.expect("Error setting Ctrl-C handler");
});
Ok(())
}
impl<'cb> BpfScheduler<'cb> {
#[allow(clippy::too_many_arguments)]
pub fn init(
open_object: &'cb mut MaybeUninit<OpenObject>,
open_opts: Option<bpf_object_open_opts>,
exit_dump_len: u32,
partial: bool,
debug: bool,
builtin_idle: bool,
numa_local: bool,
slice_ns: u64,
name: &str,
) -> Result<Self> {
let shutdown = Arc::new(AtomicBool::new(false));
set_ctrlc_handler(shutdown.clone()).context("Error setting Ctrl-C handler")?;
let mut skel_builder = BpfSkelBuilder::default();
skel_builder.obj_builder.debug(debug);
let mut skel = scx_ops_open!(skel_builder, open_object, rustland, open_opts)?;
fn callback(data: &[u8]) -> i32 {
#[allow(static_mut_refs)]
unsafe {
BUF.0.copy_from_slice(data);
}
0
}
let topo = Topology::new().unwrap();
skel.maps.rodata_data.as_mut().unwrap().smt_enabled = topo.smt_enabled;
skel.struct_ops.rustland_mut().flags =
*compat::SCX_OPS_ENQ_LAST | *compat::SCX_OPS_ALLOW_QUEUED_WAKEUP;
if partial {
skel.struct_ops.rustland_mut().flags |= *compat::SCX_OPS_SWITCH_PARTIAL;
}
if numa_local {
skel.struct_ops.rustland_mut().flags |= *compat::SCX_OPS_BUILTIN_IDLE_PER_NODE;
}
skel.struct_ops.rustland_mut().exit_dump_len = exit_dump_len;
skel.maps.rodata_data.as_mut().unwrap().usersched_pid = std::process::id();
skel.maps.rodata_data.as_mut().unwrap().khugepaged_pid = Self::khugepaged_pid();
skel.maps.rodata_data.as_mut().unwrap().builtin_idle = builtin_idle;
skel.maps.rodata_data.as_mut().unwrap().numa_local = numa_local;
skel.maps.rodata_data.as_mut().unwrap().slice_ns = slice_ns;
skel.maps.rodata_data.as_mut().unwrap().debug = debug;
let _ = Self::set_scx_ops_name(&mut skel.struct_ops.rustland_mut().name, name);
let mut skel = scx_ops_load!(skel, rustland, uei)?;
let struct_ops = Some(scx_ops_attach!(skel, rustland)?);
let maps = &skel.maps;
let queued_ring_buffer = &maps.queued;
let mut rbb = libbpf_rs::RingBufferBuilder::new();
rbb.add(queued_ring_buffer, callback)
.expect("failed to add ringbuf callback");
let queued = rbb.build().expect("failed to build ringbuf");
let dispatched = libbpf_rs::UserRingBuffer::new(&maps.dispatched)
.expect("failed to create user ringbuf");
ALLOCATOR.lock_memory();
ALLOCATOR.disable_mmap().expect("Failed to disable mmap");
if partial {
let err = Self::use_sched_ext();
if err < 0 {
return Err(anyhow::Error::msg(format!(
"sched_setscheduler error: {err}"
)));
}
}
Ok(Self {
skel,
shutdown,
queued,
dispatched,
struct_ops,
})
}
fn set_scx_ops_name(name_field: &mut [i8], src: &str) -> Result<()> {
if !src.is_ascii() {
bail!("name must be an ASCII string");
}
let bytes = src.as_bytes();
let n = bytes.len().min(name_field.len().saturating_sub(1));
name_field.fill(0);
for i in 0..n {
name_field[i] = bytes[i] as i8;
}
let version_suffix = ::scx_utils::build_id::ops_version_suffix(env!("CARGO_PKG_VERSION"));
let bytes = version_suffix.as_bytes();
let mut i = 0;
let mut bytes_idx = 0;
let mut found_null = false;
while i < name_field.len() - 1 {
found_null |= name_field[i] == 0;
if !found_null {
i += 1;
continue;
}
if bytes_idx < bytes.len() {
name_field[i] = bytes[bytes_idx] as i8;
bytes_idx += 1;
} else {
break;
}
i += 1;
}
name_field[i] = 0;
Ok(())
}
fn khugepaged_pid() -> u32 {
let procs = match all_processes() {
Ok(p) => p,
Err(_) => return 0,
};
for proc in procs {
let proc = match proc {
Ok(p) => p,
Err(_) => continue,
};
if let Ok(stat) = proc.stat() {
if proc.exe().is_err() && stat.comm == "khugepaged" {
return proc.pid() as u32;
}
}
}
0
}
pub fn notify_complete(&mut self, nr_pending: u64) {
self.skel.maps.bss_data.as_mut().unwrap().nr_scheduled = nr_pending;
std::thread::yield_now();
}
#[allow(dead_code)]
pub fn nr_online_cpus_mut(&mut self) -> &mut u64 {
&mut self.skel.maps.bss_data.as_mut().unwrap().nr_online_cpus
}
#[allow(dead_code)]
pub fn nr_running_mut(&mut self) -> &mut u64 {
&mut self.skel.maps.bss_data.as_mut().unwrap().nr_running
}
#[allow(dead_code)]
pub fn nr_queued_mut(&mut self) -> &mut u64 {
&mut self.skel.maps.bss_data.as_mut().unwrap().nr_queued
}
#[allow(dead_code)]
pub fn nr_scheduled_mut(&mut self) -> &mut u64 {
&mut self.skel.maps.bss_data.as_mut().unwrap().nr_scheduled
}
#[allow(dead_code)]
pub fn nr_user_dispatches_mut(&mut self) -> &mut u64 {
&mut self.skel.maps.bss_data.as_mut().unwrap().nr_user_dispatches
}
#[allow(dead_code)]
pub fn nr_kernel_dispatches_mut(&mut self) -> &mut u64 {
&mut self
.skel
.maps
.bss_data
.as_mut()
.unwrap()
.nr_kernel_dispatches
}
#[allow(dead_code)]
pub fn nr_cancel_dispatches_mut(&mut self) -> &mut u64 {
&mut self
.skel
.maps
.bss_data
.as_mut()
.unwrap()
.nr_cancel_dispatches
}
#[allow(dead_code)]
pub fn nr_bounce_dispatches_mut(&mut self) -> &mut u64 {
&mut self
.skel
.maps
.bss_data
.as_mut()
.unwrap()
.nr_bounce_dispatches
}
#[allow(dead_code)]
pub fn nr_failed_dispatches_mut(&mut self) -> &mut u64 {
&mut self
.skel
.maps
.bss_data
.as_mut()
.unwrap()
.nr_failed_dispatches
}
#[allow(dead_code)]
pub fn nr_sched_congested_mut(&mut self) -> &mut u64 {
&mut self.skel.maps.bss_data.as_mut().unwrap().nr_sched_congested
}
fn use_sched_ext() -> i32 {
#[cfg(target_env = "gnu")]
let param: sched_param = sched_param { sched_priority: 0 };
#[cfg(target_env = "musl")]
let param: sched_param = sched_param {
sched_priority: 0,
sched_ss_low_priority: 0,
sched_ss_repl_period: timespec {
tv_sec: 0,
tv_nsec: 0,
},
sched_ss_init_budget: timespec {
tv_sec: 0,
tv_nsec: 0,
},
sched_ss_max_repl: 0,
};
unsafe { pthread_setschedparam(pthread_self(), SCHED_EXT, ¶m as *const sched_param) }
}
#[allow(dead_code)]
pub fn select_cpu(&mut self, pid: i32, cpu: i32, flags: u64) -> i32 {
let prog = &mut self.skel.progs.rs_select_cpu;
let mut args = task_cpu_arg {
pid: pid as c_int,
cpu: cpu as c_int,
flags: flags as c_ulong,
};
let input = ProgramInput {
context_in: Some(unsafe {
std::slice::from_raw_parts_mut(
&mut args as *mut _ as *mut u8,
std::mem::size_of_val(&args),
)
}),
..Default::default()
};
let out = prog.test_run(input).unwrap();
out.return_value as i32
}
#[allow(static_mut_refs)]
pub fn dequeue_task(&mut self) -> Result<Option<QueuedTask>, i32> {
let bss_data = self.skel.maps.bss_data.as_mut().unwrap();
match self.queued.consume_raw_n(1) {
0 => {
bss_data.nr_queued = 0;
Ok(None)
}
1 => {
let task = unsafe { EnqueuedMessage::from_bytes(&BUF.0).to_queued_task() };
bss_data.nr_queued = bss_data.nr_queued.saturating_sub(1);
Ok(Some(task))
}
res if res < 0 => Err(res),
res => panic!("Unexpected return value from libbpf-rs::consume_raw(): {res}"),
}
}
pub fn dispatch_task(&mut self, task: &DispatchedTask) -> Result<(), libbpf_rs::Error> {
let mut urb_sample = self
.dispatched
.reserve(std::mem::size_of::<bpf_intf::dispatched_task_ctx>())?;
let bytes = urb_sample.as_mut();
let dispatched_task = plain::from_mut_bytes::<bpf_intf::dispatched_task_ctx>(bytes)
.expect("failed to convert bytes");
let bpf_intf::dispatched_task_ctx {
pid,
cpu,
flags,
slice_ns,
vtime,
enq_cnt,
..
} = dispatched_task;
*pid = task.pid;
*cpu = task.cpu;
*flags = task.flags;
*slice_ns = task.slice_ns;
*vtime = task.vtime;
*enq_cnt = task.enq_cnt;
self.dispatched
.submit(urb_sample)
.expect("failed to submit task");
Ok(())
}
pub fn exited(&mut self) -> bool {
self.shutdown.load(Ordering::Relaxed) || uei_exited!(&self.skel, uei)
}
pub fn shutdown_and_report(&mut self) -> Result<UserExitInfo> {
let _ = self.struct_ops.take();
uei_report!(&self.skel, uei)
}
}
impl Drop for BpfScheduler<'_> {
fn drop(&mut self) {
if let Some(struct_ops) = self.struct_ops.take() {
drop(struct_ops);
}
ALLOCATOR.unlock_memory();
}
}