Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 4 additions & 21 deletions ffi.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@

ffibuilder = cffi.FFI()

# NOTE: no Py_BEGIN_ALLOW_THREADS wrappers are needed here; CFFI API-mode
# calls already release the GIL around C function calls.
source_code = """
#include <fcntl.h>
#include <linux/poll.h>
Expand All @@ -13,25 +15,7 @@
#include <sys/uio.h>
#include <netinet/in.h>
#include <unistd.h>
#include <Python.h>
#include "liburing.h"

int io_uring_wait_cqe_nogil(struct io_uring *ring, struct io_uring_cqe **cqe_ptr) {
int res;
Py_BEGIN_ALLOW_THREADS
res = io_uring_wait_cqe(ring, cqe_ptr);
Py_END_ALLOW_THREADS
return res;
}

int io_uring_wait_cqe_timeout_nogil(struct io_uring *ring, struct io_uring_cqe **cqe_ptr,
struct __kernel_timespec *ts) {
int res;
Py_BEGIN_ALLOW_THREADS
res = io_uring_wait_cqe_timeout(ring, cqe_ptr, ts);
Py_END_ALLOW_THREADS
return res;
}
"""


Expand Down Expand Up @@ -321,6 +305,8 @@
static inline void io_uring_prep_write(struct io_uring_sqe *sqe, int fd, const void *buf, unsigned nbytes, __u64 offset);

struct io_uring_sqe *io_uring_get_sqe(struct io_uring *ring);
static inline unsigned io_uring_sq_ready(const struct io_uring *ring);
static inline unsigned io_uring_sq_space_left(const struct io_uring *ring);
static inline void io_uring_sqe_set_data64(struct io_uring_sqe *sqe, __u64 data);
static inline void io_uring_sqe_set_flags(struct io_uring_sqe *sqe, unsigned flags);

Expand All @@ -329,9 +315,6 @@
static inline int io_uring_wait_cqe(struct io_uring *ring, struct io_uring_cqe **cqe_ptr);
static inline void io_uring_cqe_seen(struct io_uring *ring, struct io_uring_cqe *cqe);
int io_uring_wait_cqe_timeout(struct io_uring *ring, struct io_uring_cqe **cqe_ptr, struct __kernel_timespec *ts);

int io_uring_wait_cqe_nogil(struct io_uring *ring, struct io_uring_cqe **cqe_ptr);
int io_uring_wait_cqe_timeout_nogil(struct io_uring *ring, struct io_uring_cqe **cqe_ptr, struct __kernel_timespec *ts);
""")

# TODO: remove this trick, i feel it is not a official way.
Expand Down
2 changes: 1 addition & 1 deletion tests/e2e/proactor/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ async def init_proactor() -> AsyncGenerator[IoUringProactor, None]:
async def _run_proactor_task():
while True:
proactor._poll(timeout=0.0) # type: ignore[reportPrivateUsage]
await asyncio.sleep(0.1)
await asyncio.sleep(0.01)

task = loop.create_task(_run_proactor_task())

Expand Down
60 changes: 60 additions & 0 deletions tests/unit/test_submission.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
import errno

import pytest

import uringloop.proactor
from uringloop.proactor import _ProactorSubmit


def test_flush_retries_partial_submissions(monkeypatch):
submitter = _ProactorSubmit(object(), {})
submitter._pending_submit = True
submissions = iter((1, 1))
ready = iter((1, 0))
submit_calls = 0

def submit(ring):
nonlocal submit_calls
submit_calls += 1
return next(submissions)

monkeypatch.setattr(uringloop.proactor, "io_uring_submit", submit)
monkeypatch.setattr(uringloop.proactor, "io_uring_sq_ready", lambda ring: next(ready))

submitter.flush()

assert submit_calls == 2
assert not submitter._pending_submit


def test_flush_preserves_pending_state_on_error(monkeypatch):
submitter = _ProactorSubmit(object(), {})
submitter._pending_submit = True

def submit(ring):
raise OSError(errno.EAGAIN, "try again")

monkeypatch.setattr(uringloop.proactor, "io_uring_submit", submit)

with pytest.raises(OSError, match="try again"):
submitter.flush()

assert submitter._pending_submit


def test_linked_submission_reserves_all_slots(monkeypatch):
submitter = _ProactorSubmit(object(), {})
submitter._pending_submit = True
available = iter((1, 2))
flushed = False

def flush():
nonlocal flushed
flushed = True

monkeypatch.setattr(uringloop.proactor, "io_uring_sq_space_left", lambda ring: next(available))
monkeypatch.setattr(submitter, "flush", flush)

submitter.ensure_capacity(2)

assert flushed
11 changes: 10 additions & 1 deletion uringloop/lib.py
Original file line number Diff line number Diff line change
Expand Up @@ -430,10 +430,19 @@ def io_uring_cqe_seen(ring: IoUring, cqe: IoUringCqe) -> None:
lib.io_uring_cqe_seen(ring, cqe)


def io_uring_submit(ring: IoUring) -> None:
def io_uring_submit(ring: IoUring) -> int:
res = lib.io_uring_submit(ring)
if res < 0:
raise OSError(-res, os.strerror(-res))
return res


def io_uring_sq_ready(ring: IoUring) -> int:
return lib.io_uring_sq_ready(ring)


def io_uring_sq_space_left(ring: IoUring) -> int:
return lib.io_uring_sq_space_left(ring)


def io_uring_peek_cqe(ring: IoUring) -> IoUringCqe:
Expand Down
41 changes: 38 additions & 3 deletions uringloop/proactor.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,8 @@
io_uring_prep_write,
io_uring_queue_exit,
io_uring_queue_init,
io_uring_sq_ready,
io_uring_sq_space_left,
io_uring_sqe_set_data64,
io_uring_sqe_set_flags,
io_uring_submit,
Expand Down Expand Up @@ -85,7 +87,7 @@ class PendingCompletion:
ProatorCache: TypeAlias = dict[int, PendingCompletion]


DEFAULT_ENTRIES = 16
DEFAULT_ENTRIES = 256


class _IoUringFuture(futures.Future[Any]):
Expand Down Expand Up @@ -114,18 +116,26 @@ def __init__(self, ring: IoUring, cache: ProatorCache) -> None:
self._iouring = ring
self._cache: ProatorCache = cache
self._unsubmitted: list[tuple[int, KernelRequest]] = []
self._pending_submit = False

def _get_sqe(self) -> Any:
sqe = io_uring_get_sqe(self._iouring)
if sqe is None:
# submission queue is full: flush the prepared SQEs to the kernel
# to free up slots, then retry
io_uring_submit(self._iouring)
self.flush()
sqe = io_uring_get_sqe(self._iouring)
if sqe is None:
raise RuntimeError("io_uring submission queue is full")
return sqe

def ensure_capacity(self, count: int) -> None:
"""Ensure a group of linked SQEs fits in one submission."""
if io_uring_sq_space_left(self._iouring) < count:
self.flush()
if io_uring_sq_space_left(self._iouring) < count:
raise RuntimeError(f"io_uring submission queue cannot fit {count} linked entries")

def recv(self, request: RecvRequest, user_data: int, flags: int = 0) -> Self:
sqe = self._get_sqe()
io_uring_prep_recv(sqe, request)
Expand Down Expand Up @@ -226,10 +236,29 @@ def poll_add(self, request: PollAddRequest, user_data: int, flags: int = 0) -> S
return self

def submit(self, op: BaseOperation, fut: _IoUringFuture | None):
"""Register the prepared SQEs; the syscall is deferred to flush().

Batching all SQEs prepared between two polls into one
io_uring_submit call is what makes io_uring cheaper than epoll:
one io_uring_enter per loop iteration instead of one per operation.
"""
for user_data, request in self._unsubmitted:
self._cache[user_data] = PendingCompletion(op, fut, request)
io_uring_submit(self._iouring)
self._unsubmitted = []
self._pending_submit = True

def flush(self):
while self._pending_submit:
try:
submitted = io_uring_submit(self._iouring)
except OSError as exc:
if exc.errno == errno.EINTR:
continue
raise
if io_uring_sq_ready(self._iouring) == 0:
self._pending_submit = False
elif submitted == 0:
raise RuntimeError("io_uring made no progress submitting queued entries")


class IoUringProactor:
Expand Down Expand Up @@ -369,6 +398,8 @@ def connect(self, conn: socket.socket, address: pyAddress) -> futures.Future[Non

def sendfile(self, sock: socket.socket, file: BufferedReader, offset: int, count: int) -> futures.Future[int]:
"""NOTE: edge cases would be handled by event loop"""
# Linked SQEs cannot cross submission boundaries.
self.submitter.ensure_capacity(2)
pipe_r, pipe_w = os.pipe()

f2p_request = SpliceRequest(
Expand Down Expand Up @@ -423,6 +454,10 @@ def poll_add(self, file: socket.socket | IOBase | int, poll_mask: int) -> future
return fut

def _poll(self, timeout: float | None = None):
# push every SQE prepared since the last poll to the kernel in a
# single syscall before waiting for completions
self.submitter.flush()

if timeout is None:
while True:
try:
Expand Down
Loading