use super::*;
pub(in crate::thread) static RUNTIME_THREAD_EXTENSION_OPS: ThreadExtensionOps =
ThreadExtensionOps {
on_switch_in: runtime_thread_switch_in_hook,
on_switch_out: runtime_thread_switch_out_hook,
on_exit: runtime_thread_exit_hook,
on_deadline_overrun: runtime_thread_deadline_overrun_hook,
drop: runtime_thread_drop_hook,
};
pub(in crate::thread) unsafe fn runtime_thread_extension(data: usize) -> ThreadExtension {
let os_extension = unsafe { runtime_thread_data_from_raw(data) }
.os_extension
.as_ref();
let scheduler_tick_cpu_time = os_extension.and_then(ThreadExtension::scheduler_tick_cpu_time);
let scheduler_tick_gate = os_extension.and_then(ThreadExtension::scheduler_tick_work_gate);
let forwards_running_policy = os_extension
.and_then(ThreadExtension::running_policy_applied_hook)
.is_some();
let mut extension = unsafe { ThreadExtension::new(data, &RUNTIME_THREAD_EXTENSION_OPS) };
if let Some(accounting) = scheduler_tick_cpu_time {
extension = extension.with_scheduler_tick_cpu_time(accounting);
}
if let Some(gate) = scheduler_tick_gate {
extension =
unsafe { extension.with_scheduler_tick_work(gate, runtime_thread_scheduler_tick_hook) };
}
if forwards_running_policy {
extension = unsafe {
extension.with_running_policy_applied_hook(runtime_thread_policy_applied_hook)
};
}
extension
}
unsafe extern "Rust" fn runtime_thread_switch_in_hook(
data: usize,
thread: ThreadId,
base_policy: SchedulePolicy,
charged_runtime_ns: u64,
) {
let runtime = unsafe { runtime_thread_data_from_raw(data) };
if let Some(extension) = runtime.os_extension.as_ref() {
unsafe {
(extension.ops().on_switch_in)(
extension.data(),
thread,
base_policy,
charged_runtime_ns,
)
};
}
}
unsafe extern "Rust" fn runtime_thread_switch_out_hook(
data: usize,
thread: ThreadId,
reason: SwitchReason,
) {
let runtime = unsafe { runtime_thread_data_from_raw(data) };
if let Some(extension) = runtime.os_extension.as_ref() {
unsafe { (extension.ops().on_switch_out)(extension.data(), thread, reason) };
}
}
unsafe extern "Rust" fn runtime_thread_exit_hook(data: usize, thread: ThreadId) {
let runtime = unsafe { runtime_thread_data_from_raw(data) };
if let Some(extension) = runtime.os_extension.as_ref() {
unsafe { (extension.ops().on_exit)(extension.data(), thread) };
}
publish_runtime_exit_completion(runtime);
}
pub(super) fn publish_runtime_exit_completion(runtime: &RuntimeThreadData) {
if runtime
.exit_completed
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
runtime.join_wait.notify_all();
}
}
unsafe extern "Rust" fn runtime_thread_deadline_overrun_hook(data: usize, thread: ThreadId) {
let runtime = unsafe { runtime_thread_data_from_raw(data) };
if let Some(extension) = runtime.os_extension.as_ref() {
unsafe { (extension.ops().on_deadline_overrun)(extension.data(), thread) };
}
}
unsafe extern "Rust" fn runtime_thread_policy_applied_hook(
data: usize,
thread: ThreadId,
base_policy: SchedulePolicy,
observed_ns: u64,
) {
let runtime = unsafe { runtime_thread_data_from_raw(data) };
let Some(extension) = runtime.os_extension.as_ref() else {
panic!(
"runtime policy forwarding lost OS extension for thread {:#x}",
thread.as_u64()
);
};
if !unsafe { extension.forward_running_policy_applied(thread, base_policy, observed_ns) } {
panic!(
"runtime policy forwarding lost callback for thread {:#x}",
thread.as_u64()
);
}
}
unsafe extern "Rust" fn runtime_thread_scheduler_tick_hook(
data: usize,
thread: ThreadId,
observed_ns: u64,
) -> SchedulerTickWorkDisposition {
let runtime = unsafe { runtime_thread_data_from_raw(data) };
let Some(extension) = runtime.os_extension.as_ref() else {
panic!(
"runtime scheduler-tick forwarding lost OS extension for thread {:#x}",
thread.as_u64()
);
};
unsafe { extension.forward_scheduler_tick_work(thread, observed_ns) }.unwrap_or_else(|| {
panic!(
"runtime scheduler-tick forwarding lost callback for thread {:#x}",
thread.as_u64()
)
})
}
unsafe extern "Rust" fn runtime_thread_drop_hook(data: usize) {
drop(unsafe { Box::from_raw(ptr::with_exposed_provenance_mut::<RuntimeThreadData>(data)) });
}
unsafe fn runtime_thread_data_from_raw(data: usize) -> &'static RuntimeThreadData {
unsafe { &*ptr::with_exposed_provenance::<RuntimeThreadData>(data) }
}
pub fn thread_os_extension(
thread: &ThreadHandle,
) -> Result<Option<ThreadOsExtensionBorrow<'_>>, TaskError> {
let runtime = task_system()
.ok_or(TaskError::NotInitialized)?
.thread_extension(thread)?;
let RuntimeExtensionKind::Runtime = classify_runtime_extension(
runtime.as_ref().map(|extension| extension.ops()),
runtime.as_ref().map_or(0, |extension| extension.data()),
)?
else {
return Ok(None);
};
let Some(runtime) = runtime else {
unreachable!("classified runtime extension must be present")
};
let data = unsafe { runtime_thread_data_from_raw(runtime.data()) };
Ok(data
.os_extension
.as_ref()
.map(|extension| ThreadOsExtensionBorrow {
data: extension.data(),
ops: extension.ops(),
_runtime: runtime,
}))
}
pub fn current_os_extension() -> Result<Option<ThreadOsExtensionLease>, TaskError> {
let runtime = current_thread_extension()?;
let RuntimeExtensionKind::Runtime = classify_runtime_extension(
runtime.as_ref().map(|extension| extension.ops()),
runtime.as_ref().map_or(0, |extension| extension.data()),
)?
else {
return Ok(None);
};
let Some(runtime) = runtime else {
unreachable!("classified runtime extension must be present")
};
let data = unsafe { runtime_thread_data_from_raw(runtime.data()) };
Ok(data
.os_extension
.as_ref()
.map(|extension| ThreadOsExtensionLease {
data: extension.data(),
ops: extension.ops(),
_runtime: runtime,
}))
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(in crate::thread) enum RuntimeExtensionKind {
Missing,
Runtime,
}
pub(in crate::thread) fn classify_runtime_extension(
ops: Option<&ThreadExtensionOps>,
data: usize,
) -> Result<RuntimeExtensionKind, TaskError> {
let Some(ops) = ops else {
return Ok(RuntimeExtensionKind::Missing);
};
if !core::ptr::eq(ops, &RUNTIME_THREAD_EXTENSION_OPS) {
return Err(TaskError::InvalidConfiguration);
}
if data == 0 || !data.is_multiple_of(core::mem::align_of::<RuntimeThreadData>()) {
return Err(TaskError::InvalidRuntimeHandle);
}
Ok(RuntimeExtensionKind::Runtime)
}