nbio(windows): poll sockets through AFD

This commit is contained in:
kalsprite
2026-08-21 21:31:49 -07:00
parent 1b4e14fbc5
commit 7bb55ba45b
4 changed files with 167 additions and 99 deletions

View File

@@ -19,6 +19,38 @@ import win "core:sys/windows"
@(private="package")
_FULLY_SUPPORTED :: true
// Poll is driven by AFD, the socket driver underneath winsock.
// `WSAEventSelect` is edge triggered (`FD_WRITE` is only recorded again after a
// send fails with WOULDBLOCK) and neither `select` nor `WSAPoll` reports send
// buffer space, so neither can give the level triggered readiness `poll` promises.
// AFD also completes on the IOCP, which makes a poll an ordinary overlapped operation.
IOCTL_AFD_POLL :: 0x00012024
SIO_BASE_HANDLE :: win.DWORD(0x48000022)
AFD_POLL_RECEIVE :: 0x0001
AFD_POLL_RECEIVE_EXPEDITED :: 0x0002
AFD_POLL_SEND :: 0x0004
AFD_POLL_DISCONNECT :: 0x0008
AFD_POLL_ABORT :: 0x0010
AFD_POLL_LOCAL_CLOSE :: 0x0020
AFD_POLL_ACCEPT :: 0x0080
AFD_POLL_CONNECT_FAIL :: 0x0100
AFD_Poll_Handle_Info :: struct {
handle: win.HANDLE,
events: win.ULONG,
status: win.NTSTATUS,
}
AFD_Poll_Info :: struct {
timeout: i64,
number_of_handles: win.ULONG,
exclusive: win.ULONG,
handles: [1]AFD_Poll_Handle_Info,
}
afd_device_name := [?]u16{'\\','D','e','v','i','c','e','\\','A','f','d','\\','E','n','d','p','o','i','n','t'}
@(private="package")
_Event_Loop :: struct {
timeouts: avl.Tree(^Operation),
@@ -88,7 +120,7 @@ _Timeout :: struct {
@(private="package")
_Poll :: struct {
wait_handle: win.HANDLE,
info: AFD_Poll_Info,
}
@(private="package")
@@ -699,19 +731,6 @@ _remove :: proc(target: ^Operation) {
target._impl.timeout = (^Operation)(REMOVED)
switch target.type {
case .Poll:
win.UnregisterWaitEx(target.poll._impl.wait_handle, win.INVALID_HANDLE_VALUE)
target.poll._impl.wait_handle = nil
ok := win.PostQueuedCompletionStatus(
g.iocp,
0,
0,
&target._impl.over,
)
ensure(ok == true, "unexpected PostQueuedCompletionStatus error")
return
case .Timeout:
if avl.remove_value(&target.l.timeouts, target) {
debug("removed timeout directly")
@@ -727,6 +746,17 @@ _remove :: proc(target: ^Operation) {
// Synchronous ops, picked up in handler.
return
case .Poll:
// The poll may have completed already, with its completion queued but not yet
// handled, `NOT_FOUND` is expected rather than exceptional.
if !win.CancelIoEx(g.afd, &target._impl.over) {
#partial switch win.System_Error(win.GetLastError()) {
case .NOT_FOUND:
// nop
case: assert(false, "unexpected CancelIoEx error")
}
}
case .Accept, .Dial, .Read, .Recv, .Send, .Write, .Send_File:
if is_pending(target._impl.over) {
handle := operation_handle(target)
@@ -824,6 +854,7 @@ g: struct{
mu: sync.Mutex,
refs: int,
iocp: win.HANDLE,
afd: win.HANDLE,
err: General_Error,
}
@@ -839,6 +870,36 @@ g_ref :: proc() -> General_Error {
if g.iocp == nil {
g.err = General_Error(win.GetLastError())
}
if g.err != nil { return g.err }
// A handle on the socket driver, used to poll sockets for readiness.
iosb: win.IO_STATUS_BLOCK
status := win.NtCreateFile(
&g.afd,
win.SYNCHRONIZE,
&{
Length = size_of(win.OBJECT_ATTRIBUTES),
ObjectName = &{
Length = u16(len(afd_device_name)*2),
MaximumLength = u16(len(afd_device_name)*2),
Buffer = raw_data(afd_device_name[:]),
},
},
&iosb,
nil,
0,
win.FILE_SHARE_READ|win.FILE_SHARE_WRITE,
win.FILE_OPEN,
0,
nil,
0,
)
if syserr := win.System_Error(win.RtlNtStatusToDosError(status)); syserr != .SUCCESS {
g.err = General_Error(syserr)
} else if win.CreateIoCompletionPort(g.afd, g.iocp, 0, 0) != g.iocp {
g.err = General_Error(win.GetLastError())
}
}
sync.atomic_add(&g.refs, 1)
@@ -850,6 +911,7 @@ g_unref :: proc() {
sync.guard(&g.mu)
if sync.atomic_sub(&g.refs, 1) == 1 {
if g.afd != nil { win.CloseHandle(g.afd) }
win.CloseHandle(g.iocp)
g.err = nil
}
@@ -872,7 +934,7 @@ operation_handle :: proc(op: ^Operation) -> win.HANDLE {
case .Recv: return win.HANDLE(uintptr(net.any_socket_to_socket(op.recv.socket)))
case .Send: return win.HANDLE(uintptr(net.any_socket_to_socket(op.send.socket)))
case .Send_File: return win.HANDLE(uintptr(net.any_socket_to_socket(op.sendfile.socket)))
case .Poll: return win.HANDLE(uintptr(net.any_socket_to_socket(op.poll.socket)))
case .Poll: return g.afd
case .Stat: return win.HANDLE(uintptr(op.stat.handle))
case .Timeout, .Open, ._Splice, ._Link_Timeout, ._Remove, .None:
@@ -1549,120 +1611,94 @@ sendfile_callback :: proc(op: ^Operation) -> Op_Result {
return .Done
}
// Bit indices into `WSANETWORKEVENTS.iErrorCode`, corresponding to the `FD_*` masks.
FD_READ_BIT :: 0
FD_WRITE_BIT :: 1
@(require_results)
poll_exec :: proc(op: ^Operation) -> Op_Result {
assert(op.type == .Poll)
op._impl.over = {} // Operations are recycled, clear stale state from a previous use.
events: i32 = win.FD_CLOSE
events: win.ULONG = AFD_POLL_ABORT|AFD_POLL_DISCONNECT|AFD_POLL_LOCAL_CLOSE|AFD_POLL_CONNECT_FAIL
switch op.poll.event {
case .Send: events |= win.FD_WRITE|win.FD_CONNECT
case .Receive: events |= win.FD_READ|win.FD_ACCEPT
case .Receive: events |= AFD_POLL_RECEIVE|AFD_POLL_RECEIVE_EXPEDITED|AFD_POLL_ACCEPT
case .Send: events |= AFD_POLL_SEND
case:
op.poll.result = .Invalid_Argument
return .Done
}
op._impl.over.hEvent = win.WSACreateEvent()
if win.WSAEventSelect(
// AFD needs the socket underneath any layered service providers.
base: win.SOCKET
bytes: win.DWORD
if win.WSAIoctl(
win.SOCKET(net.any_socket_to_socket(op.poll.socket)),
op._impl.over.hEvent,
events,
SIO_BASE_HANDLE,
nil, 0,
&base, size_of(base),
&bytes, nil, nil,
) != 0 {
#partial switch win.System_Error(win.GetLastError()) {
#partial switch win.System_Error(win.WSAGetLastError()) {
case .WSAEINVAL, .WSAENOTSOCK: op.poll.result = .Invalid_Argument
case: op.poll.result = .Error
}
return .Done
}
timeout := win.INFINITE
// A negative timeout is relative, in 100ns units.
timeout := max(i64)
if op.poll.expires != {} {
diff := max(0, time.diff(op.l.now, op.poll.expires))
timeout = win.DWORD(diff / time.Millisecond)
timeout = -i64(diff / 100)
}
ok := win.RegisterWaitForSingleObject(
&op.poll._impl.wait_handle,
op._impl.over.hEvent,
wait_callback,
op,
timeout,
win.WT_EXECUTEINWAITTHREAD|win.WT_EXECUTEONLYONCE,
op.poll._impl.info = {
timeout = timeout,
number_of_handles = 1,
handles = {{handle = win.HANDLE(uintptr(base)), events = events}},
}
// The OVERLAPPED doubles as the IO_STATUS_BLOCK, their first two fields line up.
status := win.NtDeviceIoControlFile(
g.afd,
nil,
nil,
&op._impl.over,
win.PIO_STATUS_BLOCK(rawptr(&op._impl.over)),
IOCTL_AFD_POLL,
&op.poll._impl.info,
size_of(AFD_Poll_Info),
&op.poll._impl.info,
size_of(AFD_Poll_Info),
)
ensure(ok == true, "unexpected RegisterWaitForSingleObject error")
return .Pending
wait_callback :: proc "system" (lpParameter: win.PVOID, TimerOrWaitFired: win.BOOLEAN) {
op := (^Operation)(lpParameter)
assert_contextless(op.type == .Poll)
if TimerOrWaitFired {
op.poll.result = .Timeout
}
ok := win.PostQueuedCompletionStatus(
g.iocp,
0,
0,
&op._impl.over,
)
ensure_contextless(ok == true, "unexpected PostQueuedCompletionStatus error")
// The AFD handle is not set to skip completion on success, so a completion is
// queued even when this finishes synchronously.
#partial switch win.System_Error(win.RtlNtStatusToDosError(status)) {
case .SUCCESS, .IO_PENDING:
return .Pending
case:
op.poll.result = .Error
return .Done
}
}
poll_callback :: proc(op: ^Operation) {
assert(op.type == .Poll)
// Clear the socket's internal network event record, and find out what actually
// fired. Without this the record stays set after an event is reported, so the next
// `WSAEventSelect` on that socket signals its event object immediately from the
// stale record. That completes a poll for a readiness that never happened, and the
// send/recv the caller then makes fails with WOULDBLOCK.
if op._impl.over.hEvent != nil {
nev: win.WSANETWORKEVENTS
sk := win.SOCKET(net.any_socket_to_socket(op.poll.socket))
if win.WSAEnumNetworkEvents(sk, op._impl.over.hEvent, &nev) == 0 {
bit: uint
switch op.poll.event {
case .Receive: bit = FD_READ_BIT
case .Send: bit = FD_WRITE_BIT
}
// Only downgrade a result that is still `Ready`; `wait_callback` may have
// already set `.Timeout`.
if op.poll.result == nil && nev.lNetworkEvents & (i32(1) << bit) != 0 && nev.iErrorCode[bit] != 0 {
op.poll.result = .Error
}
}
}
// Tear down in the reverse order of `poll_exec`: stop `wait_callback` from running
// before the event it waits on goes away. `INVALID_HANDLE_VALUE` waits for an
// in-flight callback to return, so it can't touch `op` after it is recycled.
if op.poll._impl.wait_handle != nil {
win.UnregisterWaitEx(op.poll._impl.wait_handle, win.INVALID_HANDLE_VALUE)
op.poll._impl.wait_handle = nil
}
if op._impl.over.hEvent != nil {
win.WSACloseEvent(op._impl.over.hEvent)
op._impl.over.hEvent = nil
}
if op.poll.result != nil {
return
}
_, err := get_result(op._impl.over)
#partial switch err {
case .SUCCESS:
case:
// AFD reports a timeout by coming back with no handles.
if op.poll._impl.info.number_of_handles == 0 {
op.poll.result = .Timeout
return
}
if _, err := get_result(op._impl.over); err != .SUCCESS {
op.poll.result = .Error
return
}
if op.poll._impl.info.handles[0].events & (AFD_POLL_ABORT|AFD_POLL_CONNECT_FAIL) != 0 {
op.poll.result = .Error
}
}

View File

@@ -1626,6 +1626,9 @@ Poll a socket for readiness.
NOTE: this is provided to help with "legacy" APIs that require polling behavior.
If you can avoid it and use the other procs in this package, do so.
NOTE: on Windows only one poll per socket is delivered, a second poll on the same
socket does not complete.
Any user data can be set on the returned operation's `user_data` field.
Polymorphic variants for type safe user data are available under `poll_poly`, `poll_poly2`, and `poll_poly3`.
@@ -1656,6 +1659,9 @@ Poll a socket for readiness.
NOTE: this is provided to help with "legacy" APIs that require polling behavior.
If you can avoid it and use the other procs in this package, do so.
NOTE: on Windows only one poll per socket is delivered, a second poll on the same
socket does not complete.
This procedure uses polymorphism for type safe user data up to a certain size.
Inputs:
@@ -1690,6 +1696,9 @@ Poll a socket for readiness.
NOTE: this is provided to help with "legacy" APIs that require polling behavior.
If you can avoid it and use the other procs in this package, do so.
NOTE: on Windows only one poll per socket is delivered, a second poll on the same
socket does not complete.
This procedure uses polymorphism for type safe user data up to a certain size.
Inputs:
@@ -1725,6 +1734,9 @@ Poll a socket for readiness.
NOTE: this is provided to help with "legacy" APIs that require polling behavior.
If you can avoid it and use the other procs in this package, do so.
NOTE: on Windows only one poll per socket is delivered, a second poll on the same
socket does not complete.
This procedure uses polymorphism for type safe user data up to a certain size.
Inputs:

View File

@@ -7,6 +7,19 @@ foreign import ntdll_lib "system:ntdll.lib"
foreign ntdll_lib {
RtlGetVersion :: proc(lpVersionInformation: ^OSVERSIONINFOEXW) -> NTSTATUS ---
NtDeviceIoControlFile :: proc(
FileHandle: HANDLE,
Event: HANDLE,
ApcRoutine: PIO_APC_ROUTINE,
ApcContext: rawptr,
IoStatusBlock: PIO_STATUS_BLOCK,
IoControlCode: ULONG,
InputBuffer: rawptr,
InputBufferLength: ULONG,
OutputBuffer: rawptr,
OutputBufferLength: ULONG,
) -> NTSTATUS ---
NtQueryInformationProcess :: proc(
ProcessHandle: HANDLE,

View File

@@ -212,13 +212,18 @@ remove_multiple_poll :: proc(t: ^testing.T) {
if event_loop_guard(t) {
testing.set_fail_timeout(t, time.Minute)
sock, ep := open_next_available_local_port(t)
defer nbio.close(sock)
// Two sockets rather than two polls on one socket: only one poll per socket is
// delivered on Windows, and what this tests is removal, not that.
removed_sock, removed_ep := open_next_available_local_port(t)
defer nbio.close(removed_sock)
kept_sock, kept_ep := open_next_available_local_port(t)
defer nbio.close(kept_sock)
hit: bool
first := nbio.poll(sock, .Receive, on_poll)
nbio.poll_poly2(sock, .Receive, t, &hit, on_poll2)
first := nbio.poll(removed_sock, .Receive, on_poll)
nbio.poll_poly2(kept_sock, .Receive, t, &hit, on_poll2)
on_poll :: proc(op: ^nbio.Operation) {
log.error("shouldn't be called")
@@ -235,7 +240,9 @@ remove_multiple_poll :: proc(t: ^testing.T) {
ev(t, nbio.tick(0), nil)
nbio.dial_poly(ep, t, on_dial)
// Make both readable, the removed poll must still not fire.
nbio.dial_poly(removed_ep, t, on_dial)
nbio.dial_poly(kept_ep, t, on_dial)
on_dial :: proc(op: ^nbio.Operation, t: ^testing.T) {
ev(t, op.dial.err, nil)