use std::io;
use std::os::windows::io::RawSocket;
use std::sync::{Mutex, OnceLock, TryLockError};
use std::time::{Duration, Instant};
use windows::Win32::Foundation::HANDLE;
use windows::Win32::Networking::WinSock::{
SIO_BASE_HANDLE, SIO_BSP_HANDLE, SIO_BSP_HANDLE_POLL, SIO_BSP_HANDLE_SELECT, SOCKET,
WSAGetLastError, WSAIoctl,
};
use windows::Win32::System::IO::OVERLAPPED_ENTRY;
use super::abi::AfdPollInfo;
use super::completion_port::CompletionPort;
use super::device::{AfdDevice, finished};
use super::slots::{Completion, SlotTable, Token};
use crate::{Event, Interest};
const DEVICE_KEY: usize = 0;
const GROUP_SIZE: usize = 32;
const BATCH: usize = 256;
const DRAIN_DEADLINE: Duration = Duration::from_secs(5);
pub struct AfdPort {
port: CompletionPort,
devices: Box<[OnceLock<AfdDevice>]>,
table: SlotTable,
entries: Mutex<Box<[OVERLAPPED_ENTRY]>>,
}
unsafe impl Send for AfdPort {}
unsafe impl Sync for AfdPort {}
impl AfdPort {
pub fn new(capacity: usize) -> io::Result<Self> {
if capacity == 0 || u32::try_from(capacity).is_err() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"AFD port capacity must be between 1 and u32::MAX",
));
}
let port = Self {
port: CompletionPort::new()?,
devices: (0..capacity.div_ceil(GROUP_SIZE))
.map(|_| OnceLock::new())
.collect(),
table: SlotTable::new(capacity),
entries: Mutex::new(vec![OVERLAPPED_ENTRY::default(); BATCH].into_boxed_slice()),
};
port.open_device(0)?;
Ok(port)
}
pub fn arm(&self, socket: RawSocket, interest: Interest) -> io::Result<Token> {
if !interest.readable && !interest.writable {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"an AFD poll needs readable or writable interest",
));
}
let base = base_socket(socket)?;
let index = self.table.claim().ok_or_else(|| {
io::Error::new(io::ErrorKind::QuotaExceeded, "every AFD poll slot is armed")
})?;
if self.device(index).is_none()
&& let Err(error) = self.open_device(index)
{
self.table.unclaim(index);
return Err(error);
}
let device = self.device(index).expect("the device was opened above");
let token = self.table.publish(
index,
socket as usize,
AfdPollInfo::new(HANDLE(base as _), interest),
);
let request = self.table.request(index);
let started = unsafe { device.poll(request.info, request.status_block, request.context) };
match started {
Ok(()) => Ok(token),
Err(error) => {
self.table.abandon(index);
Err(error)
}
}
}
pub fn cancel(&self, token: Token) -> io::Result<()> {
if !self.table.begin_cancel(token) {
return Ok(());
}
let cancelled = match self.devices[token.index() / GROUP_SIZE].get() {
Some(device) => unsafe { device.cancel(self.table.status_block(token.index())) },
None => Ok(()),
};
self.table.end_cancel(token.index());
cancelled
}
pub fn poll(
&self,
timeout: Option<Duration>,
mut sink: impl FnMut(Token, io::Result<Event>),
) -> io::Result<usize> {
let mut entries = match self.entries.try_lock() {
Ok(entries) => entries,
Err(TryLockError::Poisoned(poisoned)) => poisoned.into_inner(),
Err(TryLockError::WouldBlock) => {
return Err(io::Error::new(
io::ErrorKind::WouldBlock,
"another thread is polling this AFD port",
));
}
};
let dequeued = self.port.dequeue(&mut entries, timeout)?;
let mut polls = 0;
for entry in &entries[..dequeued] {
if entry.lpOverlapped.is_null() {
continue;
}
match self.table.complete(entry.lpOverlapped) {
Completion::Foreign => {}
Completion::Cancelled => polls += 1,
Completion::Finished {
token,
status,
readiness,
} => {
polls += 1;
sink(token, finished(status, readiness));
}
}
}
Ok(polls)
}
#[must_use]
pub fn armed(&self) -> usize {
self.table.outstanding()
}
pub fn wake(&self) -> io::Result<()> {
self.port.post()
}
fn device(&self, index: usize) -> Option<&AfdDevice> {
self.devices[index / GROUP_SIZE].get()
}
fn open_device(&self, index: usize) -> io::Result<()> {
let device = AfdDevice::open(&self.port, DEVICE_KEY)?;
drop(self.devices[index / GROUP_SIZE].set(device));
Ok(())
}
}
impl Drop for AfdPort {
fn drop(&mut self) {
self.shutdown(Self::cancel, DRAIN_DEADLINE);
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum Shutdown {
Drained,
Leaked,
}
impl AfdPort {
pub(super) fn shutdown(
&mut self,
cancel: impl Fn(&Self, Token) -> io::Result<()>,
limit: Duration,
) -> Shutdown {
if !self.table.is_leaked() {
let mut refused = false;
for index in 0..self.table.len() {
if let Some(token) = self.table.armed_token(index) {
refused |= cancel(self, token).is_err();
}
}
if refused || !self.drain(limit) {
self.table.leak();
}
}
if self.table.is_leaked() {
Shutdown::Leaked
} else {
Shutdown::Drained
}
}
fn drain(&self, limit: Duration) -> bool {
let deadline = Instant::now() + limit;
while self.table.outstanding() > 0 {
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() || self.poll(Some(remaining), |_, _| {}).is_err() {
return false;
}
}
true
}
}
fn base_socket(socket: RawSocket) -> io::Result<usize> {
let mut first_error = None;
for (attempt, code) in [
SIO_BASE_HANDLE,
SIO_BSP_HANDLE_SELECT,
SIO_BSP_HANDLE_POLL,
SIO_BSP_HANDLE,
]
.into_iter()
.enumerate()
{
let mut base = SOCKET(0);
let mut returned = 0_u32;
let result = unsafe {
WSAIoctl(
SOCKET(socket as usize),
code,
None,
0,
Some((&raw mut base).cast()),
size_of::<SOCKET>() as u32,
&raw mut returned,
None,
None,
)
};
if result == 0 && (attempt == 0 || base.0 != socket as usize) {
return Ok(base.0);
}
if attempt == 0 {
first_error = Some(unsafe { WSAGetLastError() }.0);
}
}
Err(io::Error::from_raw_os_error(
first_error.unwrap_or_default(),
))
}