use alloc::{collections::VecDeque, sync::Arc, vec::Vec};
use core::{
mem::replace,
ops::Deref,
sync::atomic::{AtomicBool, AtomicUsize, Ordering},
};
use ax_tracepoint::{ExtTracePoint, TracePoint};
use super::KernelTraceAux;
use crate::{
sync::{IrqMutex, Mutex},
task::future::IrqNotify,
};
const TRACEPOINT_RECLAIM_BATCH: usize = 64;
struct TracepointSnapshotState {
current: Arc<ExtTracePoint<KernelTraceAux>>,
epoch: usize,
}
struct KernelExtTracePointState {
snapshot: IrqMutex<TracepointSnapshotState>,
readers: [AtomicUsize; 2],
update: Mutex<()>,
reclaimer: &'static TracepointReclaimer,
}
struct RetiredTracepoint {
state: Arc<KernelExtTracePointState>,
snapshot: Arc<ExtTracePoint<KernelTraceAux>>,
reader_epoch: usize,
}
#[derive(Clone)]
pub struct KernelExtTracePoint {
state: Arc<KernelExtTracePointState>,
}
struct TracepointSnapshotLease<'a> {
snapshot: Option<Arc<ExtTracePoint<KernelTraceAux>>>,
readers: &'a AtomicUsize,
reclaimer: &'static TracepointReclaimer,
}
impl Drop for TracepointSnapshotLease<'_> {
fn drop(&mut self) {
drop(self.snapshot.take());
if self.readers.fetch_sub(1, Ordering::Release) == 1 {
self.reclaimer.notify.notify_irq();
}
}
}
impl Deref for TracepointSnapshotLease<'_> {
type Target = ExtTracePoint<KernelTraceAux>;
fn deref(&self) -> &Self::Target {
self.snapshot
.as_deref()
.expect("tracepoint snapshot lease was already released")
}
}
impl KernelExtTracePoint {
pub(super) fn new(
tracepoint: ExtTracePoint<KernelTraceAux>,
reclaimer: &'static TracepointReclaimer,
) -> Self {
Self {
state: Arc::new(KernelExtTracePointState {
snapshot: IrqMutex::new(TracepointSnapshotState {
current: Arc::new(tracepoint),
epoch: 0,
}),
readers: [AtomicUsize::new(0), AtomicUsize::new(0)],
update: Mutex::new(()),
reclaimer,
}),
}
}
fn acquire_snapshot(&self) -> TracepointSnapshotLease<'_> {
let snapshot = self.state.snapshot.lock();
let readers = &self.state.readers[snapshot.epoch % self.state.readers.len()];
readers.fetch_add(1, Ordering::AcqRel);
let current = Arc::clone(&snapshot.current);
drop(snapshot);
TracepointSnapshotLease {
snapshot: Some(current),
readers,
reclaimer: self.state.reclaimer,
}
}
pub fn read<R>(&self, operation: impl FnOnce(&ExtTracePoint<KernelTraceAux>) -> R) -> R {
let snapshot = self.acquire_snapshot();
operation(&snapshot)
}
pub fn update<R>(&self, operation: impl FnOnce(&mut ExtTracePoint<KernelTraceAux>) -> R) -> R {
ax_runtime::task::thread::current::validate_blocking_context()
.expect("tracepoint updates require a preemptible task context");
let _update = self.state.update.lock();
let current = {
let snapshot = self.state.snapshot.lock();
Arc::clone(&snapshot.current)
};
let tracepoint = current.trace_point();
let was_enabled = current.has_callbacks();
let mut next = current.as_ref().clone();
let result = operation(&mut next);
let is_enabled = next.has_callbacks();
let next = Arc::new(next);
let (retired, retired_epoch) = {
if was_enabled && !is_enabled {
tracepoint.set_callback_gate(false);
super::sched::publish_runtime_gate(tracepoint, false);
}
let mut snapshot = self.state.snapshot.lock();
let retired_epoch = snapshot.epoch % self.state.readers.len();
let retired = replace(&mut snapshot.current, next);
snapshot.epoch = snapshot.epoch.wrapping_add(1);
(retired, retired_epoch)
};
if !was_enabled && is_enabled {
tracepoint.set_callback_gate(true);
super::sched::publish_runtime_gate(tracepoint, true);
}
drop(current);
if self.state.readers[retired_epoch].load(Ordering::Acquire) == 0 {
drop(retired);
} else {
self.state.reclaimer.enqueue(RetiredTracepoint {
state: Arc::clone(&self.state),
snapshot: retired,
reader_epoch: retired_epoch,
});
}
result
}
pub fn trace_point(&self) -> &'static TracePoint<KernelTraceAux> {
self.read(ExtTracePoint::trace_point)
}
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
struct TracepointReclaimDrain {
pending: bool,
runnable: bool,
}
pub(super) struct TracepointReclaimer {
queue: Mutex<VecDeque<RetiredTracepoint>>,
notify: IrqNotify,
started: AtomicBool,
}
impl TracepointReclaimer {
pub(super) const fn new() -> Self {
Self {
queue: Mutex::new(VecDeque::new()),
notify: IrqNotify::new(),
started: AtomicBool::new(false),
}
}
fn enqueue(&self, retired: RetiredTracepoint) {
self.queue.lock().push_back(retired);
self.notify.notify();
}
fn drain(&self, limit: usize) -> TracepointReclaimDrain {
let retired = {
let mut queue = self.queue.lock();
let count = limit.min(queue.len());
queue.drain(..count).collect::<Vec<_>>()
};
let mut blocked = Vec::new();
for retired in retired {
if retired.state.readers[retired.reader_epoch].load(Ordering::Acquire) == 0 {
drop(retired.snapshot);
} else {
blocked.push(retired);
}
}
let mut queue = self.queue.lock();
queue.extend(blocked);
TracepointReclaimDrain {
pending: !queue.is_empty(),
runnable: queue.iter().any(|retired| {
retired.state.readers[retired.reader_epoch].load(Ordering::Acquire) == 0
}),
}
}
pub(super) fn start_worker(&'static self) -> ax_runtime::task::thread::ThreadHandle {
if self.started.swap(true, Ordering::AcqRel) {
panic!("tracepoint reclaim worker started twice");
}
crate::task::kernel_thread_builder("tracepoint-reclaim".into())
.spawn(move || {
loop {
self.notify.wait();
loop {
let drain = self.drain(TRACEPOINT_RECLAIM_BATCH);
if !drain.pending || !drain.runnable {
break;
}
ax_runtime::task::thread::current::yield_current_cpu().unwrap_or_else(
|error| panic!("tracepoint reclaim worker failed to yield: {error}"),
);
}
}
})
.expect("failed to spawn kernel thread")
}
#[cfg(axtest)]
pub(super) fn drain_for_test(&self) -> bool {
self.drain(TRACEPOINT_RECLAIM_BATCH).pending
}
}