use core::{marker::PhantomData, mem::ManuallyDrop, pin::Pin, ptr::NonNull};
use crate::{
ContextSwitchError, CpuBindingEpoch, CpuLocalError, CpuPin, ExecutionContextHeader,
current_context,
};
#[must_use = "dropping an uncommitted switch rolls back the next CPU binding"]
pub struct PreparedContextSwitch<'switch> {
next: NonNull<ExecutionContextHeader>,
next_epoch: CpuBindingEpoch,
current_context: usize,
area: crate::CpuAreaRef,
_scope: PhantomData<&'switch mut &'switch ()>,
_not_send_or_sync: PhantomData<*mut ()>,
}
impl PreparedContextSwitch<'_> {
#[doc(hidden)]
pub const fn next_header(&self) -> NonNull<ExecutionContextHeader> {
self.next
}
#[doc(hidden)]
#[inline(always)]
pub unsafe fn commit(self) {
let prepared = ManuallyDrop::new(self);
unsafe { crate::register::commit_current_context(prepared.area, prepared.current_context) };
}
}
impl Drop for PreparedContextSwitch<'_> {
fn drop(&mut self) {
let next = unsafe { Pin::new_unchecked(self.next.as_ref()) };
if unsafe { next.unbind_cpu(self.next_epoch) }.is_err() {
panic!("prepared context-switch rollback lost the next CPU binding");
}
}
}
#[must_use = "the incoming context must withdraw the previous CPU binding"]
#[derive(Debug)]
pub struct PreviousContextBinding {
previous: NonNull<ExecutionContextHeader>,
epoch: CpuBindingEpoch,
}
impl PreviousContextBinding {
pub unsafe fn finish(
self,
previous: Pin<&ExecutionContextHeader>,
) -> Result<(), ContextSwitchError> {
if previous.as_non_null() != self.previous {
return Err(ContextSwitchError::PreviousContextMismatch);
}
unsafe { previous.unbind_cpu(self.epoch) }
}
}
pub unsafe fn prepare_context_switch<'switch>(
pin: &'switch CpuPin<'_>,
previous: Pin<&ExecutionContextHeader>,
next: Pin<&ExecutionContextHeader>,
) -> Result<(PreparedContextSwitch<'switch>, PreviousContextBinding), ContextSwitchError> {
let published = current_context(pin).map_err(|error| match error {
CpuLocalError::CurrentContextMismatch => ContextSwitchError::CurrentContextMismatch,
other => ContextSwitchError::CpuLocal(other),
})?;
if published != previous.as_non_null() {
return Err(ContextSwitchError::CurrentContextMismatch);
}
let previous_binding = previous
.cpu_binding()
.filter(|binding| binding.area == pin.area())
.ok_or(ContextSwitchError::CurrentContextMismatch)?;
let next_epoch = unsafe { next.bind_cpu(pin.area()) }?;
Ok((
PreparedContextSwitch {
next: next.as_non_null(),
next_epoch,
current_context: next.as_non_null().as_ptr() as usize,
area: pin.area(),
_scope: PhantomData,
_not_send_or_sync: PhantomData,
},
PreviousContextBinding {
previous: previous.as_non_null(),
epoch: previous_binding.epoch,
},
))
}
#[cfg(all(test, feature = "host-test"))]
mod tests {
use core::mem::MaybeUninit;
use super::*;
use crate::{
CpuAreaPrefix, CpuAreaRef, CpuIndex, install_bootstrap_context, install_cpu_area,
with_cpu_pin,
};
fn on_fresh_modeled_cpu(operation: impl FnOnce(CpuAreaRef) + Send + 'static) {
std::thread::spawn(move || {
let storage = Box::leak(Box::new(MaybeUninit::<CpuAreaPrefix>::uninit()));
let base = storage.as_mut_ptr() as usize;
storage.write(CpuAreaPrefix::initialize(CpuIndex::try_from(0).unwrap(), base).unwrap());
let area = unsafe { CpuAreaRef::from_initialized_base(base) }.unwrap();
unsafe { install_cpu_area(area) }.unwrap();
operation(area);
})
.join()
.unwrap();
}
fn context_header() -> Pin<Box<ExecutionContextHeader>> {
Box::pin(ExecutionContextHeader::new())
}
#[test]
fn abandoned_prepare_rolls_back_next_binding() {
on_fresh_modeled_cpu(|area| {
let previous = context_header();
let next = context_header();
unsafe {
with_cpu_pin(|pin| {
install_bootstrap_context(pin, previous.as_ref()).unwrap();
let (prepared, _previous_binding) =
prepare_context_switch(pin, previous.as_ref(), next.as_ref()).unwrap();
assert_eq!(current_context(pin), Ok(previous.as_ref().as_non_null()));
assert_eq!(next.cpu_area(), Some(area));
drop(prepared);
assert_eq!(current_context(pin), Ok(previous.as_ref().as_non_null()));
assert_eq!(next.cpu_area(), None);
})
}
.unwrap();
});
}
#[test]
fn prepare_reports_the_domain_mismatch_before_binding_next() {
on_fresh_modeled_cpu(|_| {
let published = context_header();
let wrong_previous = context_header();
let next = context_header();
unsafe {
with_cpu_pin(|pin| {
install_bootstrap_context(pin, published.as_ref()).unwrap();
let result =
prepare_context_switch(pin, wrong_previous.as_ref(), next.as_ref());
assert!(matches!(
result,
Err(ContextSwitchError::CurrentContextMismatch)
));
assert_eq!(next.cpu_area(), None);
})
}
.unwrap();
});
}
#[test]
fn publication_precedes_incoming_unbind() {
on_fresh_modeled_cpu(|area| {
let previous = context_header();
let next = context_header();
unsafe {
with_cpu_pin(|pin| {
install_bootstrap_context(pin, previous.as_ref()).unwrap();
let (prepared, previous_binding) =
prepare_context_switch(pin, previous.as_ref(), next.as_ref()).unwrap();
assert_eq!(current_context(pin), Ok(previous.as_ref().as_non_null()));
assert_eq!(previous.cpu_area(), Some(area));
assert_eq!(next.cpu_area(), Some(area));
prepared.commit();
assert_eq!(current_context(pin), Ok(next.as_ref().as_non_null()));
assert_eq!(previous.cpu_area(), Some(area));
previous_binding.finish(previous.as_ref()).unwrap();
assert_eq!(previous.cpu_area(), None);
assert_eq!(next.cpu_area(), Some(area));
})
}
.unwrap();
});
}
#[test]
fn stale_epoch_cannot_unbind_a_new_binding() {
on_fresh_modeled_cpu(|area| {
let header = context_header();
let stale = unsafe { header.as_ref().bind_cpu(area) }.unwrap();
unsafe { header.as_ref().unbind_cpu(stale) }.unwrap();
let current = unsafe { header.as_ref().bind_cpu(area) }.unwrap();
assert_eq!(
unsafe { header.as_ref().unbind_cpu(stale) },
Err(ContextSwitchError::StalePreviousBinding)
);
assert_eq!(header.cpu_area(), Some(area));
unsafe { header.as_ref().unbind_cpu(current) }.unwrap();
});
}
}