use std::cell::UnsafeCell;
use std::ffi::c_void;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use windows::Win32::Foundation::{NTSTATUS, STATUS_CANCELLED, STATUS_PENDING};
use windows::Win32::System::IO::{IO_STATUS_BLOCK, IO_STATUS_BLOCK_0, OVERLAPPED};
use super::abi::AfdPollInfo;
use crate::Event;
const ARMED: u64 = 1;
const CANCELLING: u64 = 2;
const COMPLETED: u64 = 4;
const REQUESTED: u64 = 8;
const FLAGS: u64 = 15;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct Token {
index: u32,
generation: u32,
}
impl Token {
fn new(index: usize, generation: u32) -> Self {
Self {
index: u32::try_from(index).expect("table capacity fits a token index"),
generation,
}
}
pub(super) fn index(self) -> usize {
self.index as usize
}
}
struct Record {
info: AfdPollInfo,
status_block: IO_STATUS_BLOCK,
socket: usize,
}
struct Slot {
word: AtomicU64,
record: UnsafeCell<Record>,
}
unsafe impl Sync for Slot {}
pub(super) enum Completion {
Foreign,
Cancelled,
Finished {
token: Token,
status: NTSTATUS,
readiness: Event,
},
}
pub(super) struct Request {
pub(super) info: *mut AfdPollInfo,
pub(super) status_block: *mut IO_STATUS_BLOCK,
pub(super) context: *const c_void,
}
pub(super) struct SlotTable {
slots: Box<[Slot]>,
claimed: Box<[AtomicU64]>,
outstanding: AtomicUsize,
}
impl SlotTable {
pub(super) fn new(capacity: usize) -> Self {
Self {
slots: (0..capacity)
.map(|_| Slot {
word: AtomicU64::new(0),
record: UnsafeCell::new(Record {
info: AfdPollInfo::idle(),
status_block: IO_STATUS_BLOCK::default(),
socket: 0,
}),
})
.collect(),
claimed: (0..capacity.div_ceil(64))
.map(|_| AtomicU64::new(0))
.collect(),
outstanding: AtomicUsize::new(0),
}
}
pub(super) fn len(&self) -> usize {
self.slots.len()
}
pub(super) fn outstanding(&self) -> usize {
self.outstanding.load(Ordering::Acquire)
}
pub(super) fn leak(&mut self) {
std::mem::forget(std::mem::take(&mut self.slots));
}
pub(super) fn is_leaked(&self) -> bool {
self.slots.is_empty()
}
pub(super) fn claim(&self) -> Option<usize> {
for (word_index, word) in self.claimed.iter().enumerate() {
let mut current = word.load(Ordering::Relaxed);
loop {
let bit = (!current).trailing_zeros() as usize;
let index = word_index * 64 + bit;
if bit == 64 || index >= self.slots.len() {
break;
}
match word.compare_exchange_weak(
current,
current | (1 << bit),
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => return Some(index),
Err(actual) => current = actual,
}
}
}
None
}
pub(super) fn publish(&self, index: usize, socket: usize, info: AfdPollInfo) -> Token {
let slot = &self.slots[index];
let generation = ((slot.word.load(Ordering::Relaxed) >> 32) as u32).wrapping_add(1);
let record = slot.record.get();
unsafe {
(&raw mut (*record).info).write(info);
(&raw mut (*record).status_block).write(IO_STATUS_BLOCK {
Anonymous: IO_STATUS_BLOCK_0 {
Status: STATUS_PENDING,
},
Information: 0,
});
(&raw mut (*record).socket).write(socket);
}
self.outstanding.fetch_add(1, Ordering::AcqRel);
slot.word
.store(u64::from(generation) << 32 | ARMED, Ordering::Release);
Token::new(index, generation)
}
pub(super) fn request(&self, index: usize) -> Request {
let slot = &self.slots[index];
let record = slot.record.get();
unsafe {
Request {
info: &raw mut (*record).info,
status_block: &raw mut (*record).status_block,
context: std::ptr::from_ref(slot).cast(),
}
}
}
pub(super) fn unclaim(&self, index: usize) {
self.claimed[index / 64].fetch_and(!(1 << (index % 64)), Ordering::Release);
}
pub(super) fn abandon(&self, index: usize) {
self.release(index);
}
pub(super) fn armed_token(&self, index: usize) -> Option<Token> {
let word = self.slots[index].word.load(Ordering::Acquire);
(word & FLAGS == ARMED).then(|| Token::new(index, (word >> 32) as u32))
}
pub(super) fn begin_cancel(&self, token: Token) -> bool {
let armed = u64::from(token.generation) << 32 | ARMED;
self.slots[token.index()]
.word
.compare_exchange(
armed,
armed | CANCELLING | REQUESTED,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
}
pub(super) fn status_block(&self, index: usize) -> *const IO_STATUS_BLOCK {
let record = self.slots[index].record.get();
unsafe { &raw const (*record).status_block }
}
pub(super) fn end_cancel(&self, index: usize) {
let prior = self.slots[index]
.word
.fetch_and(!CANCELLING, Ordering::AcqRel);
if prior & COMPLETED != 0 {
self.release(index);
}
}
pub(super) fn complete(&self, context: *mut OVERLAPPED) -> Completion {
let offset = context.addr().wrapping_sub(self.slots.as_ptr().addr());
let size = size_of::<Slot>();
if !offset.is_multiple_of(size) || offset / size >= self.slots.len() {
return Completion::Foreign;
}
let index = offset / size;
let slot = &self.slots[index];
let prior = slot.word.fetch_or(COMPLETED, Ordering::AcqRel);
debug_assert!(
prior & ARMED != 0 && prior & COMPLETED == 0,
"a packet names an armed, uncompleted slot"
);
if prior & REQUESTED != 0 {
if prior & CANCELLING == 0 {
self.release(index);
}
return Completion::Cancelled;
}
let record = slot.record.get();
let (status, readiness) = unsafe {
let block = (&raw const (*record).status_block).read();
let info = (&raw const (*record).info).read();
let socket = (&raw const (*record).socket).read();
(block.Anonymous.Status, info.readiness(socket))
};
let token = Token::new(index, (prior >> 32) as u32);
self.release(index);
if status == STATUS_CANCELLED {
Completion::Cancelled
} else {
Completion::Finished {
token,
status,
readiness,
}
}
}
fn release(&self, index: usize) {
let slot = &self.slots[index];
let generation = slot.word.load(Ordering::Relaxed) >> 32;
slot.word.store(generation << 32, Ordering::Release);
self.outstanding.fetch_sub(1, Ordering::AcqRel);
self.unclaim(index);
}
}
#[cfg(test)]
mod tests;