use alloc::{
collections::{btree_map::BTreeMap, vec_deque::VecDeque},
sync::{Arc, Weak},
vec::Vec,
};
use core::{
cmp::Ordering,
future::Future,
ops::Deref,
pin::Pin,
sync::atomic::{AtomicBool, Ordering as AtomicOrdering},
task::{Poll, Waker},
time::Duration,
};
use ax_memory_addr::VirtAddr;
use ax_task::{
current,
future::{self, block_on, interruptible},
};
use hashbrown::HashMap;
use crate::{
StarryError, StarryResult,
mm::{AddrSpace, Backend, SharedPages},
sync::{LockdepMutexExt, Mutex},
task::{AsThread, ProcessData},
};
const NESTED_WAIT_QUEUE_LOCK_SUBCLASS: u32 = 1;
pub enum FutexAccessError {
Fault,
Retry,
Operation(StarryError),
}
pub fn retry_futex_nofault<T>(
operation: impl FnMut() -> Result<T, FutexAccessError>,
fault_in: impl FnMut() -> StarryResult<()>,
) -> StarryResult<T> {
retry_futex_nofault_with(operation, fault_in, ax_task::yield_now)
}
fn retry_futex_nofault_with<T>(
mut operation: impl FnMut() -> Result<T, FutexAccessError>,
mut fault_in: impl FnMut() -> StarryResult<()>,
mut retry: impl FnMut(),
) -> StarryResult<T> {
loop {
match operation() {
Ok(value) => return Ok(value),
Err(FutexAccessError::Fault) => fault_in()?,
Err(FutexAccessError::Retry) => {}
Err(FutexAccessError::Operation(error)) => return Err(error),
}
retry();
}
}
#[derive(Default)]
pub struct WaitQueue {
inner: Mutex<WaitQueueInner>,
}
#[derive(Default)]
struct WaitQueueInner {
queue: VecDeque<Waiter>,
}
struct Waiter {
waker: Waker,
bitset: u32,
state: Arc<WaiterState>,
}
struct WaiterState {
woken: AtomicBool,
cancelled: AtomicBool,
cleanup: Mutex<Option<FutexWaitCleanup>>,
}
impl WaiterState {
fn new(cleanup: Option<FutexWaitCleanup>) -> Self {
Self {
woken: AtomicBool::new(false),
cancelled: AtomicBool::new(false),
cleanup: Mutex::new(cleanup),
}
}
fn set_cleanup_if_not_cancelled(&self, cleanup: FutexWaitCleanup) -> bool {
let mut current = self.cleanup.lock();
if self.cancelled.load(AtomicOrdering::SeqCst) {
return false;
}
*current = Some(cleanup);
true
}
fn remove_from_current_queue(state: &Arc<Self>) -> bool {
let cleanup = state.cleanup.lock().clone();
if let Some(cleanup) = cleanup {
cleanup.table.remove_waiter(cleanup.key, state);
true
} else {
false
}
}
}
struct WaitIfFuture<'a, F> {
queue: &'a WaitQueue,
bitset: u32,
cleanup: Option<FutexWaitCleanup>,
condition: Option<F>,
state: Option<Arc<WaiterState>>,
}
impl<F: FnOnce() -> Result<bool, FutexAccessError> + Unpin> Future for WaitIfFuture<'_, F> {
type Output = Result<bool, FutexAccessError>;
fn poll(self: Pin<&mut Self>, cx: &mut core::task::Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
if let Some(condition) = this.condition.take() {
let mut inner = this.queue.inner.lock();
if !condition()? {
return Poll::Ready(Ok(false));
}
let state = Arc::new(WaiterState::new(this.cleanup.clone()));
inner.queue.push_back(Waiter {
waker: cx.waker().clone(),
bitset: this.bitset,
state: state.clone(),
});
this.state = Some(state);
return Poll::Pending;
}
let Some(state) = &this.state else {
return Poll::Ready(Ok(true));
};
if state.woken.load(AtomicOrdering::SeqCst) {
this.state = None;
Poll::Ready(Ok(true))
} else {
let mut inner = this.queue.inner.lock();
if let Some(waiter) = inner
.queue
.iter_mut()
.find(|waiter| Arc::ptr_eq(&waiter.state, state))
{
waiter.waker = cx.waker().clone();
}
Poll::Pending
}
}
}
impl<F> Drop for WaitIfFuture<'_, F> {
fn drop(&mut self) {
if let Some(state) = &self.state {
state.cancelled.store(true, AtomicOrdering::SeqCst);
if !WaiterState::remove_from_current_queue(state) {
self.queue.remove_waiter(state);
}
}
}
}
#[derive(Clone)]
pub struct FutexWaitCleanup {
table: Arc<FutexTable>,
key: usize,
}
impl WaitQueue {
pub fn new() -> Self {
Self::default()
}
pub fn wait_if(
&self,
bitset: u32,
timeout: Option<Duration>,
condition: impl FnOnce() -> bool + Unpin,
) -> StarryResult<bool> {
self.wait_if_with_cleanup(bitset, timeout, None, condition)
}
pub fn wait_if_with_cleanup(
&self,
bitset: u32,
timeout: Option<Duration>,
cleanup: Option<FutexWaitCleanup>,
condition: impl FnOnce() -> bool + Unpin,
) -> StarryResult<bool> {
match self.wait_if_with_cleanup_nofault(bitset, timeout, cleanup, || Ok(condition())) {
Ok(waited) => Ok(waited),
Err(FutexAccessError::Operation(error)) => Err(error),
Err(FutexAccessError::Fault | FutexAccessError::Retry) => {
unreachable!("infallible wait condition returned a user access error")
}
}
}
pub fn wait_if_with_cleanup_nofault(
&self,
bitset: u32,
timeout: Option<Duration>,
cleanup: Option<FutexWaitCleanup>,
condition: impl FnOnce() -> Result<bool, FutexAccessError> + Unpin,
) -> Result<bool, FutexAccessError> {
let timed = block_on(interruptible(future::timeout(
timeout,
WaitIfFuture {
queue: self,
bitset,
cleanup,
condition: Some(condition),
state: None,
},
)))
.map_err(|error| FutexAccessError::Operation(error.into()))?;
timed.map_err(|error| FutexAccessError::Operation(error.into()))?
}
fn wake_locked(queue: &mut VecDeque<Waiter>, count: usize, mask: u32, wakers: &mut Vec<Waker>) {
let base = wakers.len();
queue.retain(|waiter| {
if waiter.state.cancelled.load(AtomicOrdering::SeqCst) {
false
} else if wakers.len() - base >= count || (waiter.bitset & mask) == 0 {
true
} else {
waiter.state.woken.store(true, AtomicOrdering::SeqCst);
wakers.push(waiter.waker.clone());
false
}
});
}
pub fn wake(&self, count: usize, mask: u32) -> usize {
let mut wakers = Vec::new();
{
let mut inner = self.inner.lock();
Self::wake_locked(&mut inner.queue, count, mask, &mut wakers);
}
let woke = wakers.len();
for waker in wakers {
waker.wake();
}
woke
}
pub fn wake_op(
&self,
wake_count: usize,
target: &WaitQueue,
wake2_count: usize,
condition: impl FnOnce() -> Result<bool, FutexAccessError>,
) -> Result<usize, FutexAccessError> {
let mut condition = Some(condition);
let mut wakers = Vec::new();
match core::ptr::from_ref(self).cmp(&core::ptr::from_ref(target)) {
Ordering::Less => {
let mut src = self.inner.lock();
let mut dst = target.inner.lock_nested(NESTED_WAIT_QUEUE_LOCK_SUBCLASS);
let wake_second = condition.take().expect("condition used once")()?;
Self::wake_locked(&mut src.queue, wake_count, u32::MAX, &mut wakers);
if wake_second {
Self::wake_locked(&mut dst.queue, wake2_count, u32::MAX, &mut wakers);
}
}
Ordering::Greater => {
let mut dst = target.inner.lock();
let mut src = self.inner.lock_nested(NESTED_WAIT_QUEUE_LOCK_SUBCLASS);
let wake_second = condition.take().expect("condition used once")()?;
Self::wake_locked(&mut src.queue, wake_count, u32::MAX, &mut wakers);
if wake_second {
Self::wake_locked(&mut dst.queue, wake2_count, u32::MAX, &mut wakers);
}
}
Ordering::Equal => {
let mut src = self.inner.lock();
let wake_second = condition.take().expect("condition used once")()?;
Self::wake_locked(&mut src.queue, wake_count, u32::MAX, &mut wakers);
if wake_second {
Self::wake_locked(&mut src.queue, wake2_count, u32::MAX, &mut wakers);
}
}
}
let woke = wakers.len();
for waker in wakers {
waker.wake();
}
Ok(woke)
}
fn wake_requeue_locked(
src: &mut VecDeque<Waiter>,
dst: &mut VecDeque<Waiter>,
wake_count: usize,
wake_mask: u32,
requeue_count: usize,
target_cleanup: FutexWaitCleanup,
wakers: &mut Vec<Waker>,
) -> usize {
src.retain(|waiter| !waiter.state.cancelled.load(AtomicOrdering::SeqCst));
let mut index = 0;
while index < src.len() && wakers.len() < wake_count {
if (src[index].bitset & wake_mask) == 0 {
index += 1;
continue;
}
let waiter = src.remove(index).expect("waiter index checked");
waiter.state.woken.store(true, AtomicOrdering::SeqCst);
wakers.push(waiter.waker);
}
let mut requeued = 0;
while requeued < requeue_count {
let Some(waiter) = src.pop_front() else {
break;
};
if !waiter
.state
.set_cleanup_if_not_cancelled(target_cleanup.clone())
{
continue;
}
dst.push_back(waiter);
requeued += 1;
}
wakers.len() + requeued
}
pub fn wake_requeue_if(
&self,
wake_count: usize,
wake_mask: u32,
requeue_count: usize,
target_cleanup: FutexWaitCleanup,
target: &WaitQueue,
condition: impl FnOnce() -> Result<bool, FutexAccessError>,
) -> Result<Option<usize>, FutexAccessError> {
let mut condition = Some(condition);
let mut wakers = Vec::new();
let count = match core::ptr::from_ref(self).cmp(&core::ptr::from_ref(target)) {
Ordering::Less => {
let mut src = self.inner.lock();
let mut dst = target.inner.lock_nested(NESTED_WAIT_QUEUE_LOCK_SUBCLASS);
if !condition.take().expect("condition used once")()? {
return Ok(None);
}
Self::wake_requeue_locked(
&mut src.queue,
&mut dst.queue,
wake_count,
wake_mask,
requeue_count,
target_cleanup,
&mut wakers,
)
}
Ordering::Greater => {
let mut dst = target.inner.lock();
let mut src = self.inner.lock_nested(NESTED_WAIT_QUEUE_LOCK_SUBCLASS);
if !condition.take().expect("condition used once")()? {
return Ok(None);
}
Self::wake_requeue_locked(
&mut src.queue,
&mut dst.queue,
wake_count,
wake_mask,
requeue_count,
target_cleanup,
&mut wakers,
)
}
Ordering::Equal => {
let mut src = self.inner.lock();
if !condition.take().expect("condition used once")()? {
return Ok(None);
}
src.queue
.retain(|waiter| !waiter.state.cancelled.load(AtomicOrdering::SeqCst));
let mut index = 0;
while index < src.queue.len() && wakers.len() < wake_count {
if (src.queue[index].bitset & wake_mask) == 0 {
index += 1;
continue;
}
let waiter = src.queue.remove(index).expect("waiter index checked");
waiter.state.woken.store(true, AtomicOrdering::SeqCst);
wakers.push(waiter.waker);
}
wakers.len()
}
};
for waker in wakers {
waker.wake();
}
Ok(Some(count))
}
fn remove_waiter(&self, state: &Arc<WaiterState>) -> bool {
let mut inner = self.inner.lock();
inner
.queue
.retain(|waiter| !Arc::ptr_eq(&waiter.state, state));
inner.queue.is_empty()
}
pub fn is_empty(&self) -> bool {
self.inner.lock().queue.is_empty()
}
}
pub enum FutexKey {
Private {
address: usize,
},
Shared {
offset: usize,
region: Result<Weak<SharedPages>, Weak<()>>,
},
}
#[derive(Clone, Copy)]
pub enum FutexKeyMode {
Private,
Auto,
}
impl FutexKey {
pub fn new(aspace: &AddrSpace, address: usize, mode: FutexKeyMode) -> Self {
if matches!(mode, FutexKeyMode::Auto)
&& let Some(area) = aspace.find_area(VirtAddr::from_usize(address))
{
match area.backend() {
Backend::Shared(backend) => {
return Self::Shared {
offset: address - area.start().as_usize(),
region: Ok(Arc::downgrade(backend.pages())),
};
}
Backend::File(file) => {
return Self::Shared {
offset: address - area.start().as_usize(),
region: Err(file.futex_handle()),
};
}
_ => {}
}
}
Self::Private { address }
}
pub fn new_current(address: usize, mode: FutexKeyMode) -> Self {
if matches!(mode, FutexKeyMode::Private) {
return Self::Private { address };
}
let curr = current();
let aspace_arc = curr.as_thread().proc_data.aspace();
let aspace = aspace_arc.lock();
Self::new(&aspace, address, mode)
}
pub fn new_for_process_teardown(proc_data: &ProcessData, address: usize) -> Self {
let aspace_arc = proc_data.aspace();
let Some(aspace) = aspace_arc.try_lock() else {
return Self::Private { address };
};
Self::new(&aspace, address, FutexKeyMode::Auto)
}
fn as_usize(&self) -> usize {
match self {
FutexKey::Private { address } => *address,
FutexKey::Shared { offset, .. } => *offset,
}
}
}
pub struct FutexEntry {
pub wq: WaitQueue,
}
impl FutexEntry {
fn new() -> Self {
Self {
wq: WaitQueue::new(),
}
}
}
const FUTEX_SHARDS: usize = 64;
pub struct FutexTable {
buckets: [Mutex<HashMap<usize, Arc<FutexEntry>>>; FUTEX_SHARDS],
}
impl FutexTable {
#[allow(clippy::new_without_default)]
pub fn new() -> Self {
Self {
buckets: core::array::from_fn(|_| Mutex::new(HashMap::new())),
}
}
#[inline]
fn bucket(&self, key: usize) -> &Mutex<HashMap<usize, Arc<FutexEntry>>> {
let h = (key as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15);
&self.buckets[(h >> (64 - 6)) as usize % FUTEX_SHARDS]
}
pub fn is_empty(&self) -> bool {
self.buckets.iter().all(|b| b.lock().is_empty())
}
pub fn get(&self, key: &FutexKey) -> Option<FutexGuard<'_>> {
let key = key.as_usize();
let entry = self.bucket(key).lock().get(&key).cloned()?;
Some(FutexGuard {
table: self,
key,
inner: entry,
})
}
pub fn get_or_insert(&self, key: &FutexKey) -> FutexGuard<'_> {
let key = key.as_usize();
let mut bucket = self.bucket(key).lock();
let entry = bucket
.entry(key)
.or_insert_with(|| Arc::new(FutexEntry::new()));
FutexGuard {
table: self,
key,
inner: entry.clone(),
}
}
pub fn cleanup_for(self: &Arc<Self>, key: &FutexKey) -> FutexWaitCleanup {
FutexWaitCleanup {
table: self.clone(),
key: key.as_usize(),
}
}
fn remove_waiter(&self, key: usize, state: &Arc<WaiterState>) {
let mut bucket = self.bucket(key).lock();
let should_remove = if let Some(entry) = bucket.get(&key) {
entry.wq.remove_waiter(state) && Arc::strong_count(entry) == 1
} else {
false
};
if should_remove {
bucket.remove(&key);
}
}
}
#[doc(hidden)]
pub struct FutexGuard<'a> {
table: &'a FutexTable,
key: usize,
inner: Arc<FutexEntry>,
}
impl Deref for FutexGuard<'_> {
type Target = Arc<FutexEntry>;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl Drop for FutexGuard<'_> {
fn drop(&mut self) {
let mut bucket = self.table.bucket(self.key).lock();
if Arc::strong_count(&self.inner) <= 2 && self.inner.wq.is_empty() {
bucket.remove(&self.key);
}
}
}
struct FutexTables {
map: BTreeMap<usize, Arc<FutexTable>>,
operations: usize,
}
impl FutexTables {
const fn new() -> Self {
Self {
map: BTreeMap::new(),
operations: 0,
}
}
fn get_or_insert(&mut self, key: usize) -> Arc<FutexTable> {
self.operations += 1;
if self.operations == 100 {
self.operations = 0;
self.map
.retain(|_, table| Arc::strong_count(table) > 1 || !table.is_empty());
}
self.map
.entry(key)
.or_insert_with(|| Arc::new(FutexTable::new()))
.clone()
}
}
static SHARED_FUTEX_TABLES: Mutex<FutexTables> = Mutex::new(FutexTables::new());
pub fn futex_table_for(key: &FutexKey) -> Arc<FutexTable> {
let curr = current();
futex_table_for_process(curr.as_thread().proc_data.as_ref(), key)
}
pub fn futex_table_for_process(proc_data: &ProcessData, key: &FutexKey) -> Arc<FutexTable> {
match key {
FutexKey::Private { .. } => proc_data.futex_table.clone(),
FutexKey::Shared { region, .. } => {
let ptr = match region {
Ok(pages) => Weak::as_ptr(pages) as usize,
Err(key) => Weak::as_ptr(key) as usize,
};
SHARED_FUTEX_TABLES.lock().get_or_insert(ptr)
}
}
}
#[cfg(all(test, not(axtest)))]
mod tests {
use alloc::boxed::Box;
use core::{cell::Cell, task::Context};
use super::*;
#[test]
fn nofault_failure_is_transactional() {
let wait_queue = WaitQueue::new();
let mut wait = Box::pin(WaitIfFuture {
queue: &wait_queue,
bitset: u32::MAX,
cleanup: None,
condition: Some(|| Err(FutexAccessError::Fault)),
state: None,
});
let mut context = Context::from_waker(Waker::noop());
assert!(matches!(
wait.as_mut().poll(&mut context),
Poll::Ready(Err(FutexAccessError::Fault))
));
assert!(wait_queue.is_empty());
let source = WaitQueue::new();
let target = WaitQueue::new();
let state = Arc::new(WaiterState::new(None));
source.inner.lock().queue.push_back(Waiter {
waker: Waker::noop().clone(),
bitset: u32::MAX,
state: state.clone(),
});
assert!(matches!(
source.wake_op(1, &target, 1, || Err(FutexAccessError::Fault)),
Err(FutexAccessError::Fault)
));
assert_eq!(source.inner.lock().queue.len(), 1);
assert!(!state.woken.load(AtomicOrdering::SeqCst));
let target_cleanup = FutexWaitCleanup {
table: Arc::new(FutexTable::new()),
key: 0x2000,
};
assert!(matches!(
source.wake_requeue_if(1, u32::MAX, 1, target_cleanup, &target, || {
Err(FutexAccessError::Retry)
}),
Err(FutexAccessError::Retry)
));
assert_eq!(source.inner.lock().queue.len(), 1);
assert!(target.is_empty());
assert!(!state.woken.load(AtomicOrdering::SeqCst));
let attempts = Cell::new(0);
let fault_in_unlocked = Cell::new(false);
let result = retry_futex_nofault_with(
|| {
attempts.set(attempts.get() + 1);
if attempts.get() == 1 {
source.wake_op(0, &target, 0, || Err(FutexAccessError::Fault))
} else {
source.wake_op(0, &target, 0, || Ok(false))
}
},
|| {
let source_unlocked = !unsafe { source.inner.raw() }.is_owned_by_current();
let target_unlocked = !unsafe { target.inner.raw() }.is_owned_by_current();
fault_in_unlocked.set(source_unlocked && target_unlocked);
Ok(())
},
|| {},
);
assert!(matches!(result, Ok(0)));
assert_eq!(attempts.get(), 2);
assert!(fault_in_unlocked.get());
assert_eq!(source.inner.lock().queue.len(), 1);
assert!(!state.woken.load(AtomicOrdering::SeqCst));
}
}