use core::sync::atomic::Ordering;
use super::sched;
use super::tcb::{self, NO_TASK};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum JoinError {
Stale,
SelfJoin,
AlreadyJoined,
Faulted,
}
#[no_mangle]
pub extern "C" fn rivet_task_exit_core(val_lo: usize, val_hi: usize) -> ! {
if let Some(id) = sched::current() {
if let Some(t) = tcb::get(id) {
let size = t.result_size.load(Ordering::Acquire) as usize;
let buf = unsafe { &mut *t.result_buf.get() };
if size > 0 && size <= buf.len() {
if size <= 8 {
let lo = val_lo.to_le_bytes();
let hi = val_hi.to_le_bytes();
for (i, b) in lo.iter().chain(hi.iter()).take(size).enumerate() {
buf[i] = *b;
}
} else {
let src = val_lo as *const u8;
for (i, slot) in buf.iter_mut().take(size).enumerate() {
*slot = unsafe { core::ptr::read_volatile(src.add(i)) };
}
}
}
crate::critical::enter(|| {
t.exited.store(true, Ordering::Release);
let joiner = t.joiner.swap(NO_TASK, Ordering::AcqRel);
if joiner != NO_TASK {
sched::unblock(joiner);
}
});
}
}
loop {
sched::block_current();
crate::port::arch::request_reschedule();
}
}
pub fn should_stop() -> bool {
sched::current()
.and_then(tcb::get)
.map(|t| t.stop_requested.load(Ordering::Acquire))
.unwrap_or(false)
}
pub fn join_task<T: 'static + Send>(handle: &super::TaskHandle) -> Result<T, JoinError> {
let id = handle.id as usize;
let Some(t) = tcb::get(id) else {
return Err(JoinError::Stale);
};
if t.generation.load(Ordering::Acquire) != handle.generation {
return Err(JoinError::Stale);
}
let me = sched::current().unwrap_or(NO_TASK);
if Some(id) == sched::current() {
return Err(JoinError::SelfJoin);
}
if t.joiner
.compare_exchange(NO_TASK, me, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
return Err(JoinError::AlreadyJoined);
}
loop {
let must_wait = crate::critical::enter(|| {
if t.exited.load(Ordering::Acquire) {
false
} else {
sched::block_current();
true
}
});
if !must_wait {
break;
}
crate::port::arch::request_reschedule();
}
let _ = t
.joiner
.compare_exchange(me, NO_TASK, Ordering::AcqRel, Ordering::Acquire);
let size = t.result_size.load(Ordering::Acquire) as usize;
if size != core::mem::size_of::<T>() {
return Err(JoinError::Faulted);
}
Ok(unsafe { core::ptr::read(t.result_buf.get() as *const T) })
}