use std::io;
use std::os::windows::io::{AsRawSocket, AsSocket, BorrowedSocket, OwnedSocket};
use windows_sys::Win32::Foundation::HANDLE;
use windows_sys::Win32::Networking::WinSock::{
SOCKET, WSA_INVALID_EVENT, WSA_IO_PENDING, WSABUF, WSACloseEvent, WSACreateEvent, WSAEVENT,
WSAGetOverlappedResult, WSARecv, WSASend,
};
use windows_sys::Win32::System::IO::{CancelIoEx, CreateIoCompletionPort, OVERLAPPED};
use crate::operation::payload_ptr_from_overlapped;
use crate::{Completion, CompletionPort, Issued, Operation, OperationId, Submitted};
impl CompletionPort {
pub fn associate_socket(
&self,
socket: OwnedSocket,
key: usize,
) -> io::Result<AssociatedSocket<'_>> {
let result = unsafe {
CreateIoCompletionPort(
socket.as_raw_socket() as usize as HANDLE,
self.raw(),
key,
0,
)
};
if result.is_null() {
return Err(io::Error::last_os_error());
}
Ok(AssociatedSocket {
port: self,
socket,
key,
})
}
}
#[derive(Debug)]
pub struct AssociatedSocket<'port> {
port: &'port CompletionPort,
socket: OwnedSocket,
key: usize,
}
impl<'port> AssociatedSocket<'port> {
#[must_use]
pub fn socket(&self) -> BorrowedSocket<'_> {
self.socket.as_socket()
}
#[must_use]
pub fn key(&self) -> usize {
self.key
}
#[must_use]
pub fn port(&self) -> &'port CompletionPort {
self.port
}
fn raw_socket(&self) -> SOCKET {
self.socket.as_raw_socket() as usize
}
#[track_caller]
pub fn recv(&self, len: usize) -> io::Result<SocketIo> {
let socket = self.raw_socket();
let operation = Operation::new(recv_payload(len)?);
let submitted = unsafe {
self.port.submit_with(operation, |overlapped| {
let payload = payload_ptr_from_overlapped::<SocketPayload>(overlapped);
let ret = WSARecv(
socket,
std::ptr::addr_of!((*payload).wsabuf),
1,
std::ptr::null_mut(),
std::ptr::addr_of_mut!((*payload).flags),
overlapped,
None,
);
classify_socket(ret)
})
};
finish_socket(submitted)
}
#[track_caller]
pub fn send(&self, data: Vec<u8>) -> io::Result<SocketIo> {
let socket = self.raw_socket();
let operation = Operation::new(send_payload(data)?);
let submitted = unsafe {
self.port.submit_with(operation, |overlapped| {
let payload = payload_ptr_from_overlapped::<SocketPayload>(overlapped);
let ret = WSASend(
socket,
std::ptr::addr_of!((*payload).wsabuf),
1,
std::ptr::null_mut(),
0,
overlapped,
None,
);
classify_socket(ret)
})
};
finish_socket(submitted)
}
pub fn cancel(&self, id: OperationId) -> io::Result<()> {
self.port.live_operations().cancel_if_live(id, || {
let ok = unsafe { CancelIoEx(self.raw_socket() as HANDLE, id.as_ptr()) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(())
})
}
pub fn cancel_all(&self) -> io::Result<()> {
let ok = unsafe { CancelIoEx(self.raw_socket() as HANDLE, std::ptr::null()) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
}
struct SocketPayload {
buffer: Vec<u8>,
wsabuf: WSABUF,
flags: u32,
}
unsafe impl Send for SocketPayload {}
fn recv_payload(len: usize) -> io::Result<SocketPayload> {
let wsalen = checked_len(len, "receive buffer")?;
let mut buffer = vec![0_u8; len];
let wsabuf = WSABUF {
len: wsalen,
buf: buffer.as_mut_ptr(),
};
Ok(SocketPayload {
buffer,
wsabuf,
flags: 0,
})
}
fn send_payload(mut data: Vec<u8>) -> io::Result<SocketPayload> {
let wsabuf = WSABUF {
len: checked_len(data.len(), "send buffer")?,
buf: data.as_mut_ptr(),
};
Ok(SocketPayload {
buffer: data,
wsabuf,
flags: 0,
})
}
fn classify_socket(ret: i32) -> io::Result<Issued> {
if ret == 0 {
return Ok(Issued::Pending);
}
let error = io::Error::last_os_error();
if error.raw_os_error() == Some(WSA_IO_PENDING) {
Ok(Issued::Pending)
} else {
Err(error)
}
}
fn finish_socket(submitted: Submitted<SocketPayload>) -> io::Result<SocketIo> {
match submitted {
Submitted::Pending(id) => Ok(SocketIo { id }),
Submitted::Completed { .. } => Err(io::Error::other(
"socket adapter observed a synchronous completion; the socket must not be in a \
skip-on-success completion mode",
)),
Submitted::Failed { error, .. } => Err(error),
}
}
fn checked_len(len: usize, which: &str) -> io::Result<u32> {
u32::try_from(len).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!("a {which} is limited to u32::MAX bytes; {len} does not fit"),
)
})
}
#[derive(Debug)]
pub struct SocketIo {
id: OperationId,
}
impl SocketIo {
#[must_use]
pub fn id(&self) -> OperationId {
self.id
}
pub fn claim(self, completion: &Completion) -> Result<(Vec<u8>, io::Result<usize>), Self> {
if completion.id() != Some(self.id) {
return Err(self);
}
let operation = unsafe { completion.claim::<SocketPayload>() };
let buffer = operation.into_payload().buffer;
let result = match completion.error() {
Some(error) => Err(io::Error::from_raw_os_error(
error.raw_os_error().unwrap_or_default(),
)),
None => Ok(completion.bytes_transferred() as usize),
};
Ok((buffer, result))
}
}
#[derive(Debug)]
pub struct BlockingSocket {
socket: OwnedSocket,
}
impl BlockingSocket {
#[must_use]
pub fn new(socket: OwnedSocket) -> Self {
Self { socket }
}
#[must_use]
pub fn socket(&self) -> BorrowedSocket<'_> {
self.socket.as_socket()
}
fn raw_socket(&self) -> SOCKET {
self.socket.as_raw_socket() as usize
}
pub fn recv(&self, len: usize) -> io::Result<(Vec<u8>, usize)> {
let wsalen = checked_len(len, "receive buffer")?;
let mut buffer = vec![0_u8; len];
let wsabuf = WSABUF {
len: wsalen,
buf: buffer.as_mut_ptr(),
};
let received = unsafe {
self.run(|socket, overlapped| {
let mut flags = 0_u32;
WSARecv(
socket,
&wsabuf,
1,
std::ptr::null_mut(),
&mut flags,
overlapped,
None,
)
})
}?;
buffer.truncate(received);
Ok((buffer, received))
}
pub fn send(&self, data: &[u8]) -> io::Result<usize> {
let wsabuf = WSABUF {
len: checked_len(data.len(), "send buffer")?,
buf: data.as_ptr().cast_mut(),
};
unsafe {
self.run(|socket, overlapped| {
WSASend(
socket,
&wsabuf,
1,
std::ptr::null_mut(),
0,
overlapped,
None,
)
})
}
}
unsafe fn run<F>(&self, issue: F) -> io::Result<usize>
where
F: FnOnce(SOCKET, *mut OVERLAPPED) -> i32,
{
let socket = self.raw_socket();
let event: WSAEVENT = unsafe { WSACreateEvent() };
if event == WSA_INVALID_EVENT {
return Err(io::Error::last_os_error());
}
let mut overlapped: OVERLAPPED = unsafe { std::mem::zeroed() };
overlapped.hEvent = event as HANDLE;
let ret = issue(socket, &mut overlapped);
if ret != 0 {
let error = io::Error::last_os_error();
if error.raw_os_error() != Some(WSA_IO_PENDING) {
unsafe { WSACloseEvent(event) };
return Err(error);
}
}
let mut transferred = 0_u32;
let mut flags = 0_u32;
let ok = unsafe {
WSAGetOverlappedResult(
socket,
&overlapped,
&mut transferred,
i32::from(true),
&mut flags,
)
};
let result = if ok == 0 {
Err(io::Error::last_os_error())
} else {
Ok(transferred as usize)
};
unsafe { WSACloseEvent(event) };
result
}
}
#[cfg(test)]
mod tests;