use alloc::{boxed::Box, string::String};
use core::{
ptr,
sync::atomic::{AtomicBool, Ordering},
};
use crate::{
runtime::{
RuntimeStatus,
context::{RuntimeIrqGuard, runtime_current_cpu_mut, runtime_task_system},
lock::PreemptTicketLock,
resource::{
ExecutionContextHandle, KernelContextRequest, StackHandle, StackRequest,
ThreadResources, TlsHandle,
},
task_runtime,
},
sched::{CpuSet, SchedulePolicy},
sync::WaitQueue,
thread::{
SwitchReason, TaskError, ThreadExtension, ThreadExtensionOps, ThreadHandle, ThreadId,
ThreadSpec,
},
};
pub const DEFAULT_KERNEL_THREAD_STACK_SIZE: usize = 256 * 1024;
#[derive(Debug)]
pub struct ThreadBuilder {
name: String,
stack_size: usize,
stack_alignment: usize,
guard_size: usize,
policy: SchedulePolicy,
affinity: Option<CpuSet>,
os_extension: Option<ThreadExtension>,
}
impl ThreadBuilder {
pub fn new(name: String) -> Self {
Self {
name,
stack_size: DEFAULT_KERNEL_THREAD_STACK_SIZE,
stack_alignment: 16,
guard_size: 0,
policy: SchedulePolicy::default(),
affinity: None,
os_extension: None,
}
}
pub fn stack_size(mut self, stack_size: usize) -> Self {
self.stack_size = stack_size;
self
}
pub fn stack_alignment(mut self, stack_alignment: usize) -> Self {
self.stack_alignment = stack_alignment;
self
}
pub fn guard_size(mut self, guard_size: usize) -> Self {
self.guard_size = guard_size;
self
}
pub fn policy(mut self, policy: SchedulePolicy) -> Self {
self.policy = policy;
self
}
pub fn affinity(mut self, affinity: CpuSet) -> Self {
self.affinity = Some(affinity);
self
}
pub unsafe fn extension(mut self, extension: ThreadExtension) -> Self {
self.os_extension = Some(extension);
self
}
pub fn spawn<F>(self, entry: F) -> Result<KernelThreadHandle, TaskError>
where
F: FnOnce() + Send + 'static,
{
spawn_thread(self, entry)
}
}
#[derive(Debug)]
#[must_use = "kernel threads must be joined or explicitly detached as permanent"]
pub struct KernelThreadHandle {
thread: Option<ThreadHandle>,
}
impl KernelThreadHandle {
pub fn id(&self) -> ThreadId {
self.thread
.as_ref()
.expect("kernel thread handle is consumed only by ownership methods")
.id()
}
pub fn join(mut self) -> Result<(), TaskError> {
let handle = self.thread.take().ok_or(TaskError::InvalidConfiguration)?;
if crate::thread::current::current_thread_id()? == handle.id() {
return Err(TaskError::InvalidConfiguration);
}
let data = kernel_thread_data(&handle)?;
data.join_wait
.try_wait_until(|| data.exit_completed.load(Ordering::Acquire))?;
reap_joined_thread(handle)
}
pub fn detach_permanent(mut self) {
let _thread = self.thread.take();
}
}
fn reap_joined_thread(mut handle: ThreadHandle) -> Result<(), TaskError> {
match runtime_task_system()?.reap_thread_handle(handle) {
Ok(()) => Ok(()),
Err(error)
if matches!(
error.task_error(),
TaskError::ThreadBusy | TaskError::NotExited
) =>
{
handle = error.into_retry_handle();
drop(handle);
Ok(())
}
Err(error) => Err(error.task_error()),
}
}
impl ThreadBuilder {
fn stack_request(&self) -> StackRequest {
StackRequest {
usable_size: self.stack_size,
alignment: self.stack_alignment,
guard_size: self.guard_size,
}
}
}
fn spawn_thread<F>(mut spec: ThreadBuilder, entry: F) -> Result<KernelThreadHandle, TaskError>
where
F: FnOnce() + Send + 'static,
{
validate_spec(&spec)?;
let system = runtime_task_system()?;
let resources = allocate_thread_resources(system, spec.stack_request())?;
let extension_data = Box::into_raw(Box::new(KernelThreadData::new(
entry,
core::mem::take(&mut spec.name),
spec.os_extension.take(),
)))
.expose_provenance();
let extension = unsafe { ThreadExtension::new(extension_data, &KERNEL_THREAD_OPS) };
let mut thread_spec = unsafe {
ThreadSpec::new(spec.policy)
.with_extension(extension)
.with_resources(resources)
};
if let Some(affinity) = spec.affinity.take() {
thread_spec = thread_spec.with_affinity(affinity);
}
let handle = system.create_thread(thread_spec)?;
let mut irq_guard = RuntimeIrqGuard::enter();
let result = runtime_current_cpu_mut(&mut irq_guard)
.and_then(|mut cpu| system.start_thread(cpu.as_mut(), handle.id()));
drop(irq_guard);
if let Err(error) = result {
cleanup_unstarted_thread(system, handle);
return Err(error);
}
Ok(KernelThreadHandle {
thread: Some(handle),
})
}
type KernelThreadEntry = Box<dyn FnOnce() + Send + 'static>;
struct KernelThreadData {
entry: PreemptTicketLock<Option<KernelThreadEntry>>,
join_wait: WaitQueue,
exit_completed: AtomicBool,
os_extension: Option<ThreadExtension>,
_name: String,
}
impl KernelThreadData {
fn new(
entry: impl FnOnce() + Send + 'static,
name: String,
os_extension: Option<ThreadExtension>,
) -> Self {
Self {
entry: PreemptTicketLock::new(Some(Box::new(entry))),
join_wait: WaitQueue::new(),
exit_completed: AtomicBool::new(false),
os_extension,
_name: name,
}
}
}
static KERNEL_THREAD_OPS: ThreadExtensionOps = ThreadExtensionOps {
on_switch_in: kernel_thread_switch_in,
on_switch_out: kernel_thread_switch_out,
on_exit: kernel_thread_exit,
on_deadline_overrun: kernel_thread_deadline_overrun,
drop: kernel_thread_drop,
};
unsafe extern "Rust" fn kernel_thread_switch_in(
data: usize,
thread: ThreadId,
policy: SchedulePolicy,
charged_runtime_ns: u64,
) {
let data = unsafe { kernel_thread_data_from_raw(data) };
if let Some(extension) = data.os_extension.as_ref() {
unsafe {
(extension.ops().on_switch_in)(extension.data(), thread, policy, charged_runtime_ns)
};
}
}
unsafe extern "Rust" fn kernel_thread_switch_out(
data: usize,
thread: ThreadId,
reason: SwitchReason,
) {
let data = unsafe { kernel_thread_data_from_raw(data) };
if let Some(extension) = data.os_extension.as_ref() {
unsafe { (extension.ops().on_switch_out)(extension.data(), thread, reason) };
}
}
unsafe extern "Rust" fn kernel_thread_exit(data: usize, thread: ThreadId) {
let data = unsafe { kernel_thread_data_from_raw(data) };
if let Some(extension) = data.os_extension.as_ref() {
unsafe { (extension.ops().on_exit)(extension.data(), thread) };
}
publish_kernel_thread_exit_completion(data);
}
fn publish_kernel_thread_exit_completion(data: &KernelThreadData) {
if data
.exit_completed
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
data.join_wait.notify_all();
}
}
unsafe extern "Rust" fn kernel_thread_deadline_overrun(data: usize, thread: ThreadId) {
let data = unsafe { kernel_thread_data_from_raw(data) };
if let Some(extension) = data.os_extension.as_ref() {
unsafe { (extension.ops().on_deadline_overrun)(extension.data(), thread) };
}
}
unsafe extern "Rust" fn kernel_thread_drop(data: usize) {
drop(unsafe { Box::from_raw(ptr::with_exposed_provenance_mut::<KernelThreadData>(data)) });
}
unsafe fn kernel_thread_data_from_raw(data: usize) -> &'static KernelThreadData {
unsafe { &*ptr::with_exposed_provenance::<KernelThreadData>(data) }
}
unsafe extern "C" fn kernel_thread_entry() -> ! {
if let Err(error) = unsafe {
crate::runtime::switch::finish_initial_context_switch()
} {
task_runtime::fatal_invariant(9, error_code(error));
}
let extension = crate::thread::current::current_thread_extension()
.unwrap_or_else(|error| task_runtime::fatal_invariant(10, error_code(error)))
.unwrap_or_else(|| task_runtime::fatal_invariant(11, 0));
if !core::ptr::eq(extension.ops(), &KERNEL_THREAD_OPS) {
task_runtime::fatal_invariant(12, extension.data());
}
let extension = unsafe {
extension.release_for_current_thread_entry()
};
let data_raw = extension.data();
let data = unsafe { &*ptr::with_exposed_provenance::<KernelThreadData>(data_raw) };
let Some(entry) = data.entry.lock().take() else {
task_runtime::fatal_invariant(13, data_raw);
};
entry();
let exit_permit = crate::thread::current::prepare_current_exit()
.unwrap_or_else(|error| task_runtime::fatal_invariant(15, error_code(error)));
publish_kernel_thread_exit_completion(data);
crate::thread::current::commit_current_exit(exit_permit)
}
fn validate_spec(spec: &ThreadBuilder) -> Result<(), TaskError> {
if spec.stack_size == 0 || spec.stack_alignment == 0 || !spec.stack_alignment.is_power_of_two()
{
Err(TaskError::InvalidConfiguration)
} else {
Ok(())
}
}
fn kernel_thread_data(handle: &ThreadHandle) -> Result<&KernelThreadData, TaskError> {
let extension = runtime_task_system()?
.thread_extension(handle)?
.ok_or(TaskError::InvalidConfiguration)?;
if !core::ptr::eq(extension.ops(), &KERNEL_THREAD_OPS) {
return Err(TaskError::InvalidConfiguration);
}
Ok(unsafe { &*ptr::with_exposed_provenance::<KernelThreadData>(extension.data()) })
}
fn allocate_thread_resources(
system: &crate::runtime::TaskSystem,
request: StackRequest,
) -> Result<ThreadResources, TaskError> {
let stack_result = task_runtime::allocate_stack(request);
if stack_result.status != RuntimeStatus::Success {
return Err(runtime_error(stack_result.status));
}
if stack_result.handle == 0 {
return Err(TaskError::InvalidRuntimeHandle);
}
let stack = unsafe { StackHandle::from_raw(stack_result.handle) };
let tls_result = task_runtime::allocate_kernel_tls();
let tls = match (tls_result.status, tls_result.handle) {
(RuntimeStatus::Success, 0) => {
return Err(release_partial_thread_resources(
system,
stack,
TlsHandle::NONE,
TaskError::InvalidRuntimeHandle,
));
}
(RuntimeStatus::Success, handle) => {
unsafe { TlsHandle::from_raw(handle) }
}
(RuntimeStatus::Unsupported, _) => TlsHandle::NONE,
(status, _) => {
return Err(release_partial_thread_resources(
system,
stack,
TlsHandle::NONE,
runtime_error(status),
));
}
};
let context_result = task_runtime::create_kernel_context(KernelContextRequest {
stack,
entry: kernel_thread_entry,
tls,
});
if context_result.status != RuntimeStatus::Success {
return Err(release_partial_thread_resources(
system,
stack,
tls,
runtime_error(context_result.status),
));
}
if context_result.handle == 0 {
return Err(release_partial_thread_resources(
system,
stack,
tls,
TaskError::InvalidRuntimeHandle,
));
}
Ok(unsafe {
ThreadResources::new(
ExecutionContextHandle::from_raw(context_result.handle),
stack,
tls,
crate::runtime::resource::AddressSpaceToken::NONE,
)
})
}
fn release_partial_thread_resources(
system: &crate::runtime::TaskSystem,
stack: StackHandle,
tls: TlsHandle,
creation_error: TaskError,
) -> TaskError {
let resources = unsafe {
ThreadResources::new(
ExecutionContextHandle::NONE,
stack,
tls,
crate::runtime::resource::AddressSpaceToken::NONE,
)
};
system.release_unpublished_resources(resources);
creation_error
}
fn cleanup_unstarted_thread(system: &crate::runtime::TaskSystem, handle: ThreadHandle) {
let thread = handle.id();
let _result = system.mark_exited(thread);
drop(handle);
let _result = system.reap_thread(thread);
}
const fn runtime_error(status: RuntimeStatus) -> TaskError {
TaskError::RuntimeFailure(status as u32)
}
const fn error_code(error: TaskError) -> usize {
match error {
TaskError::NotInitialized => 1,
TaskError::InvalidRuntimeHandle => 2,
TaskError::NoRunnableThread => 3,
TaskError::UnsafeContext => 4,
_ => 255,
}
}
#[cfg(test)]
mod tests {
use core::sync::atomic::{AtomicUsize, Ordering};
use super::*;
static TEST_EXTENSION_OPS: ThreadExtensionOps = ThreadExtensionOps {
on_switch_in: test_extension_switch_in,
on_switch_out: test_extension_switch_out,
on_exit: test_extension_hook,
on_deadline_overrun: test_extension_hook,
drop: test_extension_drop,
};
#[test]
fn dropping_unspawned_builder_releases_owned_extension() {
let drops = AtomicUsize::new(0);
let extension = unsafe {
ThreadExtension::new(
(&drops as *const AtomicUsize).expose_provenance(),
&TEST_EXTENSION_OPS,
)
};
let builder = unsafe {
ThreadBuilder::new(String::from("drop-test")).extension(extension)
};
drop(builder);
assert_eq!(drops.load(Ordering::Acquire), 1);
}
#[test]
fn invalid_spec_releases_extension_before_runtime_lookup() {
let drops = AtomicUsize::new(0);
let extension = unsafe {
ThreadExtension::new(
(&drops as *const AtomicUsize).expose_provenance(),
&TEST_EXTENSION_OPS,
)
};
let spec = unsafe {
ThreadBuilder::new(String::from("invalid-test"))
.stack_size(0)
.extension(extension)
};
let result = validate_spec(&spec);
drop(spec);
assert_eq!(result.unwrap_err(), TaskError::InvalidConfiguration);
assert_eq!(drops.load(Ordering::Acquire), 1);
}
unsafe extern "Rust" fn test_extension_hook(_data: usize, _thread: ThreadId) {}
unsafe extern "Rust" fn test_extension_switch_in(
_data: usize,
_thread: ThreadId,
_policy: SchedulePolicy,
_charged_runtime_ns: u64,
) {
}
unsafe extern "Rust" fn test_extension_switch_out(
_data: usize,
_thread: ThreadId,
_reason: SwitchReason,
) {
}
unsafe extern "Rust" fn test_extension_drop(data: usize) {
let drops = unsafe { &*ptr::with_exposed_provenance::<AtomicUsize>(data) };
drops.fetch_add(1, Ordering::AcqRel);
}
}