use std::cell::Cell;
use std::convert::TryFrom;
use std::io::{Error, ErrorKind, Result};
use std::sync::atomic::{AtomicUsize, Ordering};
use super::PI_ASYNC_THREAD_LOCAL_ID;
thread_local! {
static CONSUMER_THREAD: Cell<usize> = const { Cell::new(0) };
}
pub(super) struct SingleTaskOwner {
thread: AtomicUsize,
}
impl SingleTaskOwner {
pub(super) const fn new() -> Self {
Self { thread: AtomicUsize::new(0) }
}
#[inline]
pub(super) fn is_current(&self) -> bool {
CONSUMER_THREAD.try_with(|thread| {
let token = thread.get();
token != 0 && self.thread.load(Ordering::Acquire) == token
}).unwrap_or(false)
}
#[inline]
pub(super) fn bind_current(&self, runtime_id: usize) -> Result<()> {
let token = CONSUMER_THREAD.try_with(|thread| {
let cached = thread.get();
if cached != 0 { return Ok(cached); }
let value = std::thread::current().id().as_u64().get();
let token = usize::try_from(value)
.map_err(|_| Error::new(ErrorKind::Other, "线程标识不能表示为 usize"))?;
if token == 0 {
return Err(Error::new(ErrorKind::Other, "线程标识不能为零"));
}
thread.set(token);
Ok(token)
}).map_err(|_| Error::new(ErrorKind::Other, "消费者线程局部存储不可用"))??;
let owner = self.thread.load(Ordering::Acquire);
if owner == token {
return Ok(());
}
if owner != 0 {
return Err(Error::new(ErrorKind::PermissionDenied, "单线程任务池只能由首次消费者线程驱动"));
}
self.claim(token)?;
PI_ASYNC_THREAD_LOCAL_ID.try_with(|id| unsafe {
if *id.get() == usize::MAX { *id.get() = runtime_id << 32; }
}).map_err(|_| Error::new(ErrorKind::Other, "运行时线程局部存储不可用"))
}
#[inline]
fn claim(&self, token: usize) -> Result<()> {
if token == 0 {
return Err(Error::new(ErrorKind::InvalidInput, "消费者标识不能为零"));
}
match self.thread.compare_exchange(0, token, Ordering::AcqRel, Ordering::Acquire) {
Ok(_) => Ok(()),
Err(owner) if owner == token => Ok(()),
Err(_) => Err(Error::new(ErrorKind::PermissionDenied, "单线程任务池只能由首次消费者线程驱动")),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{Arc, Barrier};
#[test]
fn test_owner_local_state_boundaries() {
let owner = SingleTaskOwner::new();
assert_eq!(owner.claim(0).unwrap_err().kind(), ErrorKind::InvalidInput);
assert_eq!(owner.thread.load(Ordering::Acquire), 0);
owner.claim(usize::MAX).unwrap();
owner.claim(usize::MAX).unwrap();
assert_eq!(owner.claim(1).unwrap_err().kind(), ErrorKind::PermissionDenied);
assert_eq!(owner.thread.load(Ordering::Acquire), usize::MAX);
assert_eq!(std::mem::size_of::<SingleTaskOwner>(), std::mem::size_of::<usize>());
}
#[test]
fn test_owner_first_consumer_race_has_one_winner() {
let owner = Arc::new(SingleTaskOwner::new());
let barrier = Arc::new(Barrier::new(4));
let threads: Vec<_> = (0..4).map(|_| {
let owner = owner.clone();
let barrier = barrier.clone();
std::thread::spawn(move || {
assert!(!owner.is_current());
barrier.wait();
let result = owner.bind_current(1);
assert_eq!(result.is_ok(), owner.is_current());
if let Err(error) = &result { assert_eq!(error.kind(), ErrorKind::PermissionDenied); }
result.is_ok()
})
}).collect();
let winners = threads.into_iter().map(|thread| usize::from(thread.join().unwrap())).sum::<usize>();
assert_eq!(winners, 1);
assert!(!owner.is_current());
assert_eq!(owner.bind_current(1).unwrap_err().kind(), ErrorKind::PermissionDenied);
}
#[test]
fn test_owner_query_does_not_initialize_and_binding_preserves_legacy_id() {
std::thread::spawn(|| {
let first = SingleTaskOwner::new();
let second = SingleTaskOwner::new();
assert_eq!(CONSUMER_THREAD.with(Cell::get), 0);
assert!(!first.is_current());
assert_eq!(CONSUMER_THREAD.with(Cell::get), 0);
first.bind_current(7).unwrap();
first.bind_current(8).unwrap();
second.bind_current(9).unwrap();
assert!(first.is_current());
assert!(second.is_current());
assert_eq!(PI_ASYNC_THREAD_LOCAL_ID.with(|id| unsafe { *id.get() }), 7 << 32);
}).join().unwrap();
}
}