From 778313beaca76510a5093696ccc7a7943dfd48b6 Mon Sep 17 00:00:00 2001 From: Dang Van Nghiem Date: Fri, 15 Aug 2025 03:31:46 +0700 Subject: [PATCH 1/5] Implement UDP socket support and related transport management in the event loop --- rloop/loop.py | 51 ++++++++- src/event_loop.rs | 81 +++++++++++++- src/lib.rs | 1 + src/udp.rs | 249 ++++++++++++++++++++++++++++++++++++++++++ tests/udp/test_udp.py | 87 ++++++++------- 5 files changed, 427 insertions(+), 42 deletions(-) create mode 100644 src/udp.rs diff --git a/rloop/loop.py b/rloop/loop.py index e8bda0c..8ce2e5e 100644 --- a/rloop/loop.py +++ b/rloop/loop.py @@ -660,7 +660,56 @@ async def create_datagram_endpoint( allow_broadcast=None, sock=None, ): - raise NotImplementedError + if sock is None: + if not family and not local_addr and not remote_addr: + raise ValueError('Unexpected address family') + + af = family or socket.AF_INET + if proto == 0: + if af == socket.AF_UNIX: + proto = 0 # Unix sockets don't use IPPROTO_UDP + else: + proto = socket.IPPROTO_UDP + + # Create the socket + sock = socket.socket(af, socket.SOCK_DGRAM, proto) + try: + if reuse_address: + sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + if reuse_port: + _set_reuseport(sock) + if allow_broadcast: + sock.setsockopt(socket.SOL_SOCKET, socket.SO_BROADCAST, 1) + + sock.setblocking(False) + + if local_addr: + sock.bind(local_addr) + if remote_addr: + sock.connect(remote_addr) + + except OSError: + sock.close() + raise + else: + if not hasattr(sock, 'family') or not hasattr(sock, 'type'): + raise TypeError('sock must be a socket') + if sock.type != socket.SOCK_DGRAM: + raise ValueError('sock must be a datagram socket') + if sock.gettimeout() != 0.0: + raise ValueError('sock must be non-blocking') + + # Create the transport + transport, protocol = self._udp_conn( + (sock.fileno(), sock.family), + protocol_factory, + remote_addr + ) + + # sock is now owned by the transport, prevent close + sock.detach() + + return transport, protocol #: pipes and subprocesses methods async def connect_read_pipe(self, protocol_factory, pipe): diff --git a/src/event_loop.rs b/src/event_loop.rs index 6619e41..47836f1 100644 --- a/src/event_loop.rs +++ b/src/event_loop.rs @@ -8,7 +8,7 @@ use std::{ }; use anyhow::Result; -use mio::{Interest, Poll, Token, Waker, event, net::TcpListener}; +use mio::{Interest, Poll, Token, Waker, event, net::TcpListener, unix::SourceFd}; use pyo3::prelude::*; use crate::{ @@ -18,6 +18,7 @@ use crate::{ py::{copy_context, weakset}, server::Server, tcp::{TCPReadHandle, TCPServer, TCPServerRef, TCPTransport, TCPWriteHandle}, + udp::{UDPHandle, UDPTransport}, time::Timer, }; @@ -26,6 +27,7 @@ enum IOHandle { Signals, TCPListener(TCPListenerHandleData), TCPStream(Interest), + UDPSocket, } struct PyHandleData { @@ -73,6 +75,7 @@ pub struct EventLoop { task_factory: RwLock, tcp_lstreams: papaya::HashMap>, tcp_transports: papaya::HashMap>, + udp_transports: papaya::HashMap>, thread_id: atomic::AtomicI64, watcher_child: RwLock, #[pyo3(get)] @@ -149,6 +152,7 @@ impl EventLoop { IOHandle::Py(handle) => self.handle_io_py(py, event, handle, &mut cb_handles), IOHandle::TCPListener(handle) => self.handle_io_tcpl(py, handle, &io_handles, &mut cb_handles), IOHandle::TCPStream(_) => self.handle_io_tcps(event, &mut cb_handles), + IOHandle::UDPSocket => self.handle_io_udp(event, &mut cb_handles), IOHandle::Signals => self.handle_io_signals(py, &mut state.buf, &mut cb_handles), } } @@ -254,6 +258,14 @@ impl EventLoop { } } + #[inline] + fn handle_io_udp(&self, event: &event::Event, handles_ready: &mut VecDeque) { + let fd = event.token().0; + if event.is_readable() { + handles_ready.push_back(Box::new(UDPHandle::new(fd))); + } + } + #[inline] fn handle_io_signals(&self, py: Python, buf: &mut [u8], handles_ready: &mut VecDeque) { let mut sock_guard = self.ssock.write().unwrap(); @@ -403,6 +415,36 @@ impl EventLoop { } } + // UDP socket management methods + #[inline] + pub(crate) fn udp_socket_add(&self, fd: usize, socket: mio::net::UdpSocket) { + let token = Token(fd); + let io = self.io.lock().unwrap(); + io.registry().register(&mut SourceFd(&socket.as_raw_fd()), token, Interest::READABLE).unwrap(); + std::mem::forget(socket); // Keep socket alive + self.handles_io.pin().insert(token, IOHandle::UDPSocket); + } + + #[inline] + pub(crate) fn udp_socket_rem(&self, fd: usize) { + let token = Token(fd); + if let Some(_) = self.handles_io.pin().remove(&token) { + let io = self.io.lock().unwrap(); + let _ = io.registry().deregister(&mut SourceFd(&(fd as i32))); + } + } + + #[inline] + pub(crate) fn udp_socket_close(&self, _py: Python, fd: usize) { + self.udp_socket_rem(fd); + self.udp_transports.pin().remove(&fd); + } + + #[inline(always)] + pub(crate) fn get_udp_transport(&self, fd: usize, py: Python) -> Py { + self.udp_transports.pin().get(&fd).unwrap().clone_ref(py) + } + pub(crate) fn log_exception(&self, py: Python, ctx: LogExc) -> PyResult { let handler = self.exc_handler.read().unwrap(); handler.call1( @@ -631,6 +673,7 @@ impl EventLoop { task_factory: RwLock::new(py.None()), tcp_lstreams: papaya::HashMap::with_capacity(32), tcp_transports: papaya::HashMap::with_capacity(1024), + udp_transports: papaya::HashMap::with_capacity(1024), thread_id: atomic::AtomicI64::new(0), watcher_child: RwLock::new(py.None()), _asyncgens: weakset(py)?.unbind(), @@ -1110,6 +1153,42 @@ impl EventLoop { self.tcp_transports.pin().contains_key(&fd) } + fn _udp_conn( + pyself: Py, + py: Python, + sock: (i32, i32), + protocol_factory: PyObject, + remote_addr: Option, + ) -> PyResult<(Py, PyObject)> { + use std::net::SocketAddr; + + let parsed_remote_addr = match remote_addr { + Some(addr_obj) => { + let addr_tuple: (String, u16) = addr_obj.extract(py)?; + Some(format!("{}:{}", addr_tuple.0, addr_tuple.1).parse::() + .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid remote address: {}", e)))?) + }, + None => None, + }; + + let rself = pyself.get(); + let transport = UDPTransport::from_py(py, &pyself, sock, protocol_factory, parsed_remote_addr); + let fd = transport.fd; + let pytransport = Py::new(py, transport)?; + let proto = UDPTransport::attach(&pytransport, py)?; + rself.udp_transports.pin().insert(fd, pytransport.clone_ref(py)); + + // Get the socket for registration + let socket_fd = sock.0; + let socket = unsafe { socket2::Socket::from_raw_fd(socket_fd) }; + let _ = socket.set_nonblocking(true); + let std_socket: std::net::UdpSocket = socket.into(); + let mio_socket = mio::net::UdpSocket::from_std(std_socket); + + rself.udp_socket_add(fd, mio_socket); + Ok((pytransport, proto)) + } + fn _sig_add(&self, py: Python, sig: u8, callback: PyObject, args: PyObject, context: PyObject) { let handle = Py::new(py, CBHandle::new(callback, args, context)).unwrap(); self.sig_handlers.pin().insert(sig, handle); diff --git a/src/lib.rs b/src/lib.rs index 14ab585..0dc8db2 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -10,6 +10,7 @@ mod server; mod sock; mod tcp; mod time; +mod udp; mod utils; pub(crate) fn get_lib_version() -> &'static str { diff --git a/src/udp.rs b/src/udp.rs new file mode 100644 index 0000000..9dbb79c --- /dev/null +++ b/src/udp.rs @@ -0,0 +1,249 @@ +#[cfg(unix)] +use std::os::fd::{AsRawFd, FromRawFd}; + +use mio::net::UdpSocket; +use pyo3::{prelude::*, types::PyBytes, IntoPyObject}; +use std::{ + borrow::Cow, + cell::RefCell, + collections::HashMap, + io::ErrorKind, + net::SocketAddr, + sync::atomic, +}; + +use crate::{ + event_loop::{EventLoop, EventLoopRunState}, + handles::Handle, + sock::SocketWrapper, +}; + +struct UDPTransportState { + socket: UdpSocket, + remote_addr: Option, +} + +#[pyclass(frozen, unsendable, module = "rloop._rloop")] +pub(crate) struct UDPTransport { + pub fd: usize, + state: RefCell, + pyloop: Py, + // atomics + closing: atomic::AtomicBool, + // py protocol fields + proto: PyObject, + protom_conn_lost: PyObject, + protom_datagram_received: PyObject, + protom_error_received: PyObject, + // py extras + extra: HashMap, + sock: Py, +} + +impl UDPTransport { + fn new( + py: Python, + pyloop: Py, + socket: UdpSocket, + pyproto: Bound, + socket_family: i32, + remote_addr: Option, + ) -> Self { + let fd = socket.as_raw_fd() as usize; + let state = UDPTransportState { + socket, + remote_addr, + }; + + let protom_conn_lost = pyproto.getattr(pyo3::intern!(py, "connection_lost")).unwrap().unbind(); + let protom_datagram_received = pyproto.getattr(pyo3::intern!(py, "datagram_received")).unwrap().unbind(); + let protom_error_received = pyproto.getattr(pyo3::intern!(py, "error_received")).unwrap().unbind(); + let proto = pyproto.unbind(); + + Self { + fd, + state: RefCell::new(state), + pyloop, + closing: false.into(), + proto, + protom_conn_lost, + protom_datagram_received, + protom_error_received, + extra: HashMap::new(), + sock: SocketWrapper::from_fd(py, fd, socket_family, socket2::Type::DGRAM, 0), + } + } + + pub(crate) fn from_py(py: Python, pyloop: &Py, pysock: (i32, i32), proto_factory: PyObject, remote_addr: Option) -> Self { + let sock = unsafe { socket2::Socket::from_raw_fd(pysock.0) }; + _ = sock.set_nonblocking(true); + let std_socket: std::net::UdpSocket = sock.into(); + let socket = UdpSocket::from_std(std_socket); + + let proto = proto_factory.bind(py).call0().unwrap(); + + Self::new(py, pyloop.clone_ref(py), socket, proto, pysock.1, remote_addr) + } + + pub(crate) fn attach(pyself: &Py, py: Python) -> PyResult { + let rself = pyself.borrow(py); + rself + .proto + .call_method1(py, pyo3::intern!(py, "connection_made"), (pyself.clone_ref(py),))?; + Ok(rself.proto.clone_ref(py)) + } + + #[inline] + fn call_conn_lost(&self, py: Python, exc: Option) { + _ = self.protom_conn_lost.call1(py, (exc,)); + } + + #[inline] + fn call_datagram_received(&self, py: Python, data: &[u8], addr: SocketAddr) { + let py_data = PyBytes::new(py, data); + let py_addr = (addr.ip().to_string(), addr.port()).into_pyobject(py).unwrap(); + _ = self.protom_datagram_received.call1(py, (py_data, py_addr)); + } + + #[inline] + fn call_error_received(&self, py: Python, exc: PyErr) { + _ = self.protom_error_received.call1(py, (exc,)); + } +} + +#[pymethods] +impl UDPTransport { + #[pyo3(signature = (name, default = None))] + fn get_extra_info(&self, py: Python, name: &str, default: Option) -> Option { + match name { + "socket" => Some(self.sock.clone_ref(py).into_any()), + "sockname" => self.sock.call_method0(py, pyo3::intern!(py, "getsockname")).ok(), + "peername" => { + if self.state.borrow().remote_addr.is_some() { + self.sock.call_method0(py, pyo3::intern!(py, "getpeername")).ok() + } else { + default + } + }, + _ => self.extra.get(name).map(|v| v.clone_ref(py)).or(default), + } + } + + fn is_closing(&self) -> bool { + self.closing.load(atomic::Ordering::Relaxed) + } + + fn close(&self, py: Python) { + if self + .closing + .compare_exchange(false, true, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) + .is_err() + { + return; + } + + let event_loop = self.pyloop.get(); + event_loop.udp_socket_rem(self.fd); + self.call_conn_lost(py, None); + } + + fn abort(&self, py: Python) { + self.close(py); + } + + fn set_protocol(&self, _protocol: PyObject) -> PyResult<()> { + Err(pyo3::exceptions::PyNotImplementedError::new_err( + "UDPTransport protocol cannot be changed", + )) + } + + fn get_protocol(&self, py: Python) -> PyObject { + self.proto.clone_ref(py) + } + + fn sendto(&self, py: Python, data: Cow<[u8]>, addr: Option) -> PyResult<()> { + if self.closing.load(atomic::Ordering::Relaxed) { + return Err(pyo3::exceptions::PyRuntimeError::new_err("Cannot send on closing transport")); + } + + let target_addr = match addr { + Some(addr_obj) => { + // Parse the address from Python tuple (host, port) + let addr_tuple: (String, u16) = addr_obj.extract(py)?; + Some(format!("{}:{}", addr_tuple.0, addr_tuple.1).parse::() + .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid address: {}", e)))?) + }, + None => { + // Get remote addr without holding the borrow longer than necessary + let remote_addr = self.state.borrow().remote_addr; + remote_addr + }, + }; + + match target_addr { + Some(addr) => { + // Temporarily borrow just for the send operation + match self.state.borrow().socket.send_to(&data, addr) { + Ok(_) => Ok(()), + Err(err) if err.kind() == ErrorKind::WouldBlock => { + // For UDP, we don't buffer writes like TCP - just drop the packet or return error + Err(pyo3::exceptions::PyBlockingIOError::new_err("Socket would block")) + }, + Err(err) => Err(pyo3::exceptions::PyOSError::new_err(err.to_string())), + } + }, + None => Err(pyo3::exceptions::PyValueError::new_err("No remote address specified")), + } + } +} + +pub(crate) struct UDPHandle { + fd: usize, +} + +impl UDPHandle { + pub(crate) fn new(fd: usize) -> Self { + Self { fd } + } +} + +impl Handle for UDPHandle { + fn run(&self, py: Python, event_loop: &EventLoop, _state: &mut EventLoopRunState) { + let pytransport = event_loop.get_udp_transport(self.fd, py); + let transport = pytransport.borrow(py); + + // Read datagrams from the socket + let mut buf = [0u8; 65536]; // Max UDP packet size + + loop { + // Limit the scope of the borrow + let recv_result = { + let state = transport.state.borrow(); + state.socket.recv_from(&mut buf) + }; + + match recv_result { + Ok((size, addr)) => { + // Call the protocol's datagram_received method + // Now state is not borrowed, so sendto can work + transport.call_datagram_received(py, &buf[..size], addr); + }, + Err(err) if err.kind() == ErrorKind::WouldBlock => { + // No more data available + break; + }, + Err(err) if err.kind() == ErrorKind::Interrupted => { + // Interrupted by signal, continue + continue; + }, + Err(err) => { + // Other error - call error_received and close + let py_err = pyo3::exceptions::PyOSError::new_err(err.to_string()); + transport.call_error_received(py, py_err); + event_loop.udp_socket_close(py, self.fd); + break; + }, + } + } + } +} diff --git a/tests/udp/test_udp.py b/tests/udp/test_udp.py index b6d076c..5f47874 100644 --- a/tests/udp/test_udp.py +++ b/tests/udp/test_udp.py @@ -1,9 +1,7 @@ import asyncio +import socket - -# import socket - -# import pytest +import pytest class DatagramProto(asyncio.DatagramProtocol): @@ -38,39 +36,48 @@ def connection_lost(self, exc): self.done.set_result(None) -# def test_create_datagram_endpoint_sock(loop): -# sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) -# sock.bind(('127.0.0.1', 0)) -# fut = loop.create_datagram_endpoint( -# lambda: DatagramProto(create_future=True, loop=loop), -# sock=sock) -# transport, protocol = loop.run_until_complete(fut) -# transport.close() -# loop.run_until_complete(protocol.done) -# assert protocol.state == 'CLOSED' - - -# @pytest.mark.skipif(not hasattr(socket, 'AF_UNIX'), reason='no UDS') -# def test_create_datagram_endpoint_sock_unix(loop): -# fut = loop.create_datagram_endpoint( -# lambda: DatagramProto(create_future=True, loop=loop), -# family=socket.AF_UNIX) -# transport, protocol = loop.run_until_complete(fut) -# assert transport._sock.family == socket.AF_UNIX -# transport.close() -# loop.run_until_complete(protocol.done) -# assert protocol.state == 'CLOSED' - - -# def test_create_datagram_endpoint_existing_sock_unix(loop): -# with _unix_socket_path() as path: -# sock = socket.socket(socket.AF_UNIX, type=socket.SOCK_DGRAM) -# sock.bind(path) -# sock.close() - -# coro = loop.create_datagram_endpoint( -# lambda: DatagramProto(create_future=True, loop=loop), -# path, family=socket.AF_UNIX) -# transport, protocol = loop.run_until_complete(coro) -# transport.close() -# loop.run_until_complete(protocol.done) +def test_create_datagram_endpoint_sock(loop): + sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + sock.bind(('127.0.0.1', 0)) + sock.setblocking(False) # Make socket non-blocking + fut = loop.create_datagram_endpoint( + lambda: DatagramProto(create_future=True, loop=loop), + sock=sock) + transport, protocol = loop.run_until_complete(fut) + transport.close() + loop.run_until_complete(protocol.done) + assert protocol.state == 'CLOSED' + + +@pytest.mark.skipif(not hasattr(socket, 'AF_UNIX'), reason='no UDS') +def test_create_datagram_endpoint_sock_unix(loop): + fut = loop.create_datagram_endpoint( + lambda: DatagramProto(create_future=True, loop=loop), + family=socket.AF_UNIX) + transport, protocol = loop.run_until_complete(fut) + # Check that the socket family is AF_UNIX using get_extra_info + sock_info = transport.get_extra_info('socket') + assert sock_info.family == socket.AF_UNIX + transport.close() + loop.run_until_complete(protocol.done) + assert protocol.state == 'CLOSED' + + +def test_create_datagram_endpoint_local_addr(loop): + fut = loop.create_datagram_endpoint( + lambda: DatagramProto(create_future=True, loop=loop), + local_addr=('127.0.0.1', 0)) + transport, protocol = loop.run_until_complete(fut) + transport.close() + loop.run_until_complete(protocol.done) + assert protocol.state == 'CLOSED' + + +def test_create_datagram_endpoint_remote_addr(loop): + fut = loop.create_datagram_endpoint( + lambda: DatagramProto(create_future=True, loop=loop), + remote_addr=('127.0.0.1', 12345)) + transport, protocol = loop.run_until_complete(fut) + transport.close() + loop.run_until_complete(protocol.done) + assert protocol.state == 'CLOSED' From 340f7421ce35f7aa1307c7d3c5335f94f4a8ef70 Mon Sep 17 00:00:00 2001 From: Dang Van Nghiem Date: Fri, 15 Aug 2025 08:35:24 +0700 Subject: [PATCH 2/5] Refactor UDP socket connection and streamline datagram endpoint tests --- rloop/loop.py | 6 +--- src/event_loop.rs | 24 ++++++++++------ src/udp.rs | 66 ++++++++++++++++++++++--------------------- tests/udp/test_udp.py | 16 ++++------- 4 files changed, 56 insertions(+), 56 deletions(-) diff --git a/rloop/loop.py b/rloop/loop.py index 8ce2e5e..1862b94 100644 --- a/rloop/loop.py +++ b/rloop/loop.py @@ -700,11 +700,7 @@ async def create_datagram_endpoint( raise ValueError('sock must be non-blocking') # Create the transport - transport, protocol = self._udp_conn( - (sock.fileno(), sock.family), - protocol_factory, - remote_addr - ) + transport, protocol = self._udp_conn((sock.fileno(), sock.family), protocol_factory, remote_addr) # sock is now owned by the transport, prevent close sock.detach() diff --git a/src/event_loop.rs b/src/event_loop.rs index 47836f1..2282f94 100644 --- a/src/event_loop.rs +++ b/src/event_loop.rs @@ -18,8 +18,8 @@ use crate::{ py::{copy_context, weakset}, server::Server, tcp::{TCPReadHandle, TCPServer, TCPServerRef, TCPTransport, TCPWriteHandle}, - udp::{UDPHandle, UDPTransport}, time::Timer, + udp::{UDPHandle, UDPTransport}, }; enum IOHandle { @@ -420,7 +420,9 @@ impl EventLoop { pub(crate) fn udp_socket_add(&self, fd: usize, socket: mio::net::UdpSocket) { let token = Token(fd); let io = self.io.lock().unwrap(); - io.registry().register(&mut SourceFd(&socket.as_raw_fd()), token, Interest::READABLE).unwrap(); + io.registry() + .register(&mut SourceFd(&socket.as_raw_fd()), token, Interest::READABLE) + .unwrap(); std::mem::forget(socket); // Keep socket alive self.handles_io.pin().insert(token, IOHandle::UDPSocket); } @@ -428,8 +430,9 @@ impl EventLoop { #[inline] pub(crate) fn udp_socket_rem(&self, fd: usize) { let token = Token(fd); - if let Some(_) = self.handles_io.pin().remove(&token) { + if self.handles_io.pin().remove(&token).is_some() { let io = self.io.lock().unwrap(); + #[allow(clippy::cast_possible_wrap)] let _ = io.registry().deregister(&mut SourceFd(&(fd as i32))); } } @@ -1161,13 +1164,16 @@ impl EventLoop { remote_addr: Option, ) -> PyResult<(Py, PyObject)> { use std::net::SocketAddr; - + let parsed_remote_addr = match remote_addr { Some(addr_obj) => { let addr_tuple: (String, u16) = addr_obj.extract(py)?; - Some(format!("{}:{}", addr_tuple.0, addr_tuple.1).parse::() - .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid remote address: {}", e)))?) - }, + Some( + format!("{}:{}", addr_tuple.0, addr_tuple.1) + .parse::() + .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid remote address: {e}")))?, + ) + } None => None, }; @@ -1177,14 +1183,14 @@ impl EventLoop { let pytransport = Py::new(py, transport)?; let proto = UDPTransport::attach(&pytransport, py)?; rself.udp_transports.pin().insert(fd, pytransport.clone_ref(py)); - + // Get the socket for registration let socket_fd = sock.0; let socket = unsafe { socket2::Socket::from_raw_fd(socket_fd) }; let _ = socket.set_nonblocking(true); let std_socket: std::net::UdpSocket = socket.into(); let mio_socket = mio::net::UdpSocket::from_std(std_socket); - + rself.udp_socket_add(fd, mio_socket); Ok((pytransport, proto)) } diff --git a/src/udp.rs b/src/udp.rs index 9dbb79c..5de0eed 100644 --- a/src/udp.rs +++ b/src/udp.rs @@ -2,15 +2,8 @@ use std::os::fd::{AsRawFd, FromRawFd}; use mio::net::UdpSocket; -use pyo3::{prelude::*, types::PyBytes, IntoPyObject}; -use std::{ - borrow::Cow, - cell::RefCell, - collections::HashMap, - io::ErrorKind, - net::SocketAddr, - sync::atomic, -}; +use pyo3::{IntoPyObject, prelude::*, types::PyBytes}; +use std::{borrow::Cow, cell::RefCell, collections::HashMap, io::ErrorKind, net::SocketAddr, sync::atomic}; use crate::{ event_loop::{EventLoop, EventLoopRunState}, @@ -50,13 +43,13 @@ impl UDPTransport { remote_addr: Option, ) -> Self { let fd = socket.as_raw_fd() as usize; - let state = UDPTransportState { - socket, - remote_addr, - }; + let state = UDPTransportState { socket, remote_addr }; let protom_conn_lost = pyproto.getattr(pyo3::intern!(py, "connection_lost")).unwrap().unbind(); - let protom_datagram_received = pyproto.getattr(pyo3::intern!(py, "datagram_received")).unwrap().unbind(); + let protom_datagram_received = pyproto + .getattr(pyo3::intern!(py, "datagram_received")) + .unwrap() + .unbind(); let protom_error_received = pyproto.getattr(pyo3::intern!(py, "error_received")).unwrap().unbind(); let proto = pyproto.unbind(); @@ -74,7 +67,13 @@ impl UDPTransport { } } - pub(crate) fn from_py(py: Python, pyloop: &Py, pysock: (i32, i32), proto_factory: PyObject, remote_addr: Option) -> Self { + pub(crate) fn from_py( + py: Python, + pyloop: &Py, + pysock: (i32, i32), + proto_factory: PyObject, + remote_addr: Option, + ) -> Self { let sock = unsafe { socket2::Socket::from_raw_fd(pysock.0) }; _ = sock.set_nonblocking(true); let std_socket: std::net::UdpSocket = sock.into(); @@ -124,7 +123,7 @@ impl UDPTransport { } else { default } - }, + } _ => self.extra.get(name).map(|v| v.clone_ref(py)).or(default), } } @@ -163,21 +162,25 @@ impl UDPTransport { fn sendto(&self, py: Python, data: Cow<[u8]>, addr: Option) -> PyResult<()> { if self.closing.load(atomic::Ordering::Relaxed) { - return Err(pyo3::exceptions::PyRuntimeError::new_err("Cannot send on closing transport")); + return Err(pyo3::exceptions::PyRuntimeError::new_err( + "Cannot send on closing transport", + )); } let target_addr = match addr { Some(addr_obj) => { // Parse the address from Python tuple (host, port) let addr_tuple: (String, u16) = addr_obj.extract(py)?; - Some(format!("{}:{}", addr_tuple.0, addr_tuple.1).parse::() - .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid address: {}", e)))?) - }, + Some( + format!("{}:{}", addr_tuple.0, addr_tuple.1) + .parse::() + .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid address: {e}")))?, + ) + } None => { // Get remote addr without holding the borrow longer than necessary - let remote_addr = self.state.borrow().remote_addr; - remote_addr - }, + self.state.borrow().remote_addr + } }; match target_addr { @@ -188,10 +191,10 @@ impl UDPTransport { Err(err) if err.kind() == ErrorKind::WouldBlock => { // For UDP, we don't buffer writes like TCP - just drop the packet or return error Err(pyo3::exceptions::PyBlockingIOError::new_err("Socket would block")) - }, + } Err(err) => Err(pyo3::exceptions::PyOSError::new_err(err.to_string())), } - }, + } None => Err(pyo3::exceptions::PyValueError::new_err("No remote address specified")), } } @@ -213,7 +216,7 @@ impl Handle for UDPHandle { let transport = pytransport.borrow(py); // Read datagrams from the socket - let mut buf = [0u8; 65536]; // Max UDP packet size + let mut buf = vec![0u8; 65536].into_boxed_slice(); // Max UDP packet size loop { // Limit the scope of the borrow @@ -221,28 +224,27 @@ impl Handle for UDPHandle { let state = transport.state.borrow(); state.socket.recv_from(&mut buf) }; - + match recv_result { Ok((size, addr)) => { // Call the protocol's datagram_received method // Now state is not borrowed, so sendto can work transport.call_datagram_received(py, &buf[..size], addr); - }, + } Err(err) if err.kind() == ErrorKind::WouldBlock => { // No more data available break; - }, + } Err(err) if err.kind() == ErrorKind::Interrupted => { // Interrupted by signal, continue - continue; - }, + } Err(err) => { // Other error - call error_received and close let py_err = pyo3::exceptions::PyOSError::new_err(err.to_string()); transport.call_error_received(py, py_err); event_loop.udp_socket_close(py, self.fd); break; - }, + } } } } diff --git a/tests/udp/test_udp.py b/tests/udp/test_udp.py index 5f47874..61fe2f2 100644 --- a/tests/udp/test_udp.py +++ b/tests/udp/test_udp.py @@ -40,9 +40,7 @@ def test_create_datagram_endpoint_sock(loop): sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) sock.bind(('127.0.0.1', 0)) sock.setblocking(False) # Make socket non-blocking - fut = loop.create_datagram_endpoint( - lambda: DatagramProto(create_future=True, loop=loop), - sock=sock) + fut = loop.create_datagram_endpoint(lambda: DatagramProto(create_future=True, loop=loop), sock=sock) transport, protocol = loop.run_until_complete(fut) transport.close() loop.run_until_complete(protocol.done) @@ -51,9 +49,7 @@ def test_create_datagram_endpoint_sock(loop): @pytest.mark.skipif(not hasattr(socket, 'AF_UNIX'), reason='no UDS') def test_create_datagram_endpoint_sock_unix(loop): - fut = loop.create_datagram_endpoint( - lambda: DatagramProto(create_future=True, loop=loop), - family=socket.AF_UNIX) + fut = loop.create_datagram_endpoint(lambda: DatagramProto(create_future=True, loop=loop), family=socket.AF_UNIX) transport, protocol = loop.run_until_complete(fut) # Check that the socket family is AF_UNIX using get_extra_info sock_info = transport.get_extra_info('socket') @@ -65,8 +61,8 @@ def test_create_datagram_endpoint_sock_unix(loop): def test_create_datagram_endpoint_local_addr(loop): fut = loop.create_datagram_endpoint( - lambda: DatagramProto(create_future=True, loop=loop), - local_addr=('127.0.0.1', 0)) + lambda: DatagramProto(create_future=True, loop=loop), local_addr=('127.0.0.1', 0) + ) transport, protocol = loop.run_until_complete(fut) transport.close() loop.run_until_complete(protocol.done) @@ -75,8 +71,8 @@ def test_create_datagram_endpoint_local_addr(loop): def test_create_datagram_endpoint_remote_addr(loop): fut = loop.create_datagram_endpoint( - lambda: DatagramProto(create_future=True, loop=loop), - remote_addr=('127.0.0.1', 12345)) + lambda: DatagramProto(create_future=True, loop=loop), remote_addr=('127.0.0.1', 12345) + ) transport, protocol = loop.run_until_complete(fut) transport.close() loop.run_until_complete(protocol.done) From 743baa9d3ac3312554c7868213ce58e3271c6992 Mon Sep 17 00:00:00 2001 From: Giovanni Barillari Date: Wed, 3 Sep 2025 16:40:00 +0200 Subject: [PATCH 3/5] Refactor UDP impl --- rloop/loop.py | 88 +++++++++++++++++++++++++++++------------------ src/event_loop.rs | 48 +++++++------------------- src/io.rs | 23 ++++++++----- src/udp.rs | 81 +++++++++++++++++++++---------------------- 4 files changed, 122 insertions(+), 118 deletions(-) diff --git a/rloop/loop.py b/rloop/loop.py index 1862b94..e584057 100644 --- a/rloop/loop.py +++ b/rloop/loop.py @@ -655,56 +655,78 @@ async def create_datagram_endpoint( family=0, proto=0, flags=0, - reuse_address=None, + #: not in stdlib + # reuse_address=None, reuse_port=None, allow_broadcast=None, sock=None, ): - if sock is None: - if not family and not local_addr and not remote_addr: - raise ValueError('Unexpected address family') - - af = family or socket.AF_INET - if proto == 0: - if af == socket.AF_UNIX: - proto = 0 # Unix sockets don't use IPPROTO_UDP - else: - proto = socket.IPPROTO_UDP + if sock is not None: + if getattr(sock, 'type', None) != socket.SOCK_DGRAM: + raise ValueError(f'A datagram socket was expected, got {sock!r}') + if any((local_addr, remote_addr, family, proto, flags, reuse_port, allow_broadcast)): + raise ValueError('socket modifier keyword arguments can not be used when sock is specified.') + sock.setblocking(False) + r_addr = None + else: + if not (local_addr or remote_addr): + if family == 0: + raise ValueError('unexpected address family') + addr_info = (family, proto, None, None) + elif hasattr(socket, 'AF_UNIX') and family == socket.AF_UNIX: + for addr in (local_addr, remote_addr): + if addr is not None and not isinstance(addr, str): + raise TypeError('string is expected') + addr_info = (family, proto, local_addr, remote_addr) + else: + addr_info, infos = None, None + for addr in (local_addr, remote_addr): + if addr is None: + continue + if not (isinstance(addr, tuple) and len(addr) == 2): + raise TypeError('2-tuple is expected') + infos = await self._ensure_resolved( + addr, family=family, type=socket.SOCK_DGRAM, proto=proto, flags=flags + ) + break - # Create the socket - sock = socket.socket(af, socket.SOCK_DGRAM, proto) + if not infos: + raise OSError('getaddrinfo() returned empty list') + if local_addr is not None: + addr_info = (infos[0], infos[2], infos[4], None) + if remote_addr is not None: + addr_info = (infos[0], infos[2], None, infos[4]) + if not addr_info: + raise ValueError('can not get address information') + + sock = None + r_addr = None + sfam, spro, sladdr, sraddr = addr_info try: - if reuse_address: - sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + sock = socket.socket(family=sfam, type=socket.SOCK_DGRAM, proto=spro) + #: not in stdlib + # if reuse_address: + # sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) if reuse_port: _set_reuseport(sock) if allow_broadcast: sock.setsockopt(socket.SOL_SOCKET, socket.SO_BROADCAST, 1) - sock.setblocking(False) - - if local_addr: - sock.bind(local_addr) - if remote_addr: - sock.connect(remote_addr) - + if sladdr: + sock.bind(sladdr) + if sraddr: + if not allow_broadcast: + await self.sock_connect(sock, sraddr) + r_addr = sraddr except OSError: - sock.close() + if sock is not None: + sock.close() raise - else: - if not hasattr(sock, 'family') or not hasattr(sock, 'type'): - raise TypeError('sock must be a socket') - if sock.type != socket.SOCK_DGRAM: - raise ValueError('sock must be a datagram socket') - if sock.gettimeout() != 0.0: - raise ValueError('sock must be non-blocking') # Create the transport - transport, protocol = self._udp_conn((sock.fileno(), sock.family), protocol_factory, remote_addr) - + transport, protocol = self._udp_conn((sock.fileno(), sock.family), protocol_factory, r_addr) # sock is now owned by the transport, prevent close sock.detach() - return transport, protocol #: pipes and subprocesses methods diff --git a/src/event_loop.rs b/src/event_loop.rs index 2282f94..3455698 100644 --- a/src/event_loop.rs +++ b/src/event_loop.rs @@ -8,7 +8,9 @@ use std::{ }; use anyhow::Result; -use mio::{Interest, Poll, Token, Waker, event, net::TcpListener, unix::SourceFd}; +#[cfg(unix)] +use mio::unix::SourceFd; +use mio::{Interest, Poll, Token, Waker, event, net::TcpListener}; use pyo3::prelude::*; use crate::{ @@ -19,7 +21,7 @@ use crate::{ server::Server, tcp::{TCPReadHandle, TCPServer, TCPServerRef, TCPTransport, TCPWriteHandle}, time::Timer, - udp::{UDPHandle, UDPTransport}, + udp::{UDPReadHandle, UDPTransport}, }; enum IOHandle { @@ -262,7 +264,7 @@ impl EventLoop { fn handle_io_udp(&self, event: &event::Event, handles_ready: &mut VecDeque) { let fd = event.token().0; if event.is_readable() { - handles_ready.push_back(Box::new(UDPHandle::new(fd))); + handles_ready.push_back(Box::new(UDPReadHandle::new(fd))); } } @@ -415,15 +417,13 @@ impl EventLoop { } } - // UDP socket management methods #[inline] - pub(crate) fn udp_socket_add(&self, fd: usize, socket: mio::net::UdpSocket) { + pub(crate) fn udp_socket_add(&self, fd: usize) { let token = Token(fd); - let io = self.io.lock().unwrap(); - io.registry() - .register(&mut SourceFd(&socket.as_raw_fd()), token, Interest::READABLE) - .unwrap(); - std::mem::forget(socket); // Keep socket alive + #[allow(clippy::cast_possible_wrap)] + let mut source = Source::UDPSocket(fd as i32); + let guard_poll = self.io.lock().unwrap(); + let _ = guard_poll.registry().register(&mut source, token, Interest::READABLE); self.handles_io.pin().insert(token, IOHandle::UDPSocket); } @@ -1161,37 +1161,15 @@ impl EventLoop { py: Python, sock: (i32, i32), protocol_factory: PyObject, - remote_addr: Option, + remote_addr: Option<(String, u16)>, ) -> PyResult<(Py, PyObject)> { - use std::net::SocketAddr; - - let parsed_remote_addr = match remote_addr { - Some(addr_obj) => { - let addr_tuple: (String, u16) = addr_obj.extract(py)?; - Some( - format!("{}:{}", addr_tuple.0, addr_tuple.1) - .parse::() - .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid remote address: {e}")))?, - ) - } - None => None, - }; - let rself = pyself.get(); - let transport = UDPTransport::from_py(py, &pyself, sock, protocol_factory, parsed_remote_addr); + let transport = UDPTransport::from_py(py, &pyself, sock, protocol_factory, remote_addr); let fd = transport.fd; let pytransport = Py::new(py, transport)?; let proto = UDPTransport::attach(&pytransport, py)?; rself.udp_transports.pin().insert(fd, pytransport.clone_ref(py)); - - // Get the socket for registration - let socket_fd = sock.0; - let socket = unsafe { socket2::Socket::from_raw_fd(socket_fd) }; - let _ = socket.set_nonblocking(true); - let std_socket: std::net::UdpSocket = socket.into(); - let mio_socket = mio::net::UdpSocket::from_std(std_socket); - - rself.udp_socket_add(fd, mio_socket); + rself.udp_socket_add(fd); Ok((pytransport, proto)) } diff --git a/src/io.rs b/src/io.rs index a0cd255..1329975 100644 --- a/src/io.rs +++ b/src/io.rs @@ -8,15 +8,19 @@ use std::os::windows::io::RawSocket; use mio::{Interest, Registry, Token, event::Source as MioSource, net::TcpListener}; pub(crate) enum Source { + #[cfg(unix)] + FD(RawFd), + #[cfg(windows)] + FD(RawSocket), TCPListener(TcpListener), #[cfg(unix)] TCPStream(RawFd), #[cfg(windows)] TCPStream(RawSocket), #[cfg(unix)] - FD(RawFd), + UDPSocket(RawFd), #[cfg(windows)] - FD(RawSocket), + UDPSocket(RawSocket), } #[cfg(windows)] @@ -43,36 +47,39 @@ impl MioSource for Source { #[inline] fn register(&mut self, registry: &Registry, token: Token, interests: Interest) -> std::io::Result<()> { match self { - Self::TCPListener(inner) => inner.register(registry, token, interests), - Self::TCPStream(inner) => SourceFd(inner).register(registry, token, interests), #[cfg(unix)] Self::FD(inner) => SourceFd(inner).register(registry, token, interests), #[cfg(windows)] Self::FD(inner) => SourceRawSocket(inner).register(registry, token, interests), + Self::TCPListener(inner) => inner.register(registry, token, interests), + Self::TCPStream(inner) => SourceFd(inner).register(registry, token, interests), + Self::UDPSocket(inner) => SourceFd(inner).register(registry, token, interests), } } #[inline] fn reregister(&mut self, registry: &Registry, token: Token, interests: Interest) -> std::io::Result<()> { match self { - Self::TCPListener(inner) => inner.reregister(registry, token, interests), - Self::TCPStream(inner) => SourceFd(inner).reregister(registry, token, interests), #[cfg(unix)] Self::FD(inner) => SourceFd(inner).reregister(registry, token, interests), #[cfg(windows)] Self::FD(inner) => SourceRawSocket(inner).register(registry, token, interests), + Self::TCPListener(inner) => inner.reregister(registry, token, interests), + Self::TCPStream(inner) => SourceFd(inner).reregister(registry, token, interests), + Self::UDPSocket(inner) => SourceFd(inner).reregister(registry, token, interests), } } #[inline] fn deregister(&mut self, registry: &Registry) -> std::io::Result<()> { match self { - Self::TCPListener(inner) => inner.deregister(registry), - Self::TCPStream(inner) => SourceFd(inner).deregister(registry), #[cfg(unix)] Self::FD(inner) => SourceFd(inner).deregister(registry), #[cfg(windows)] Self::FD(inner) => SourceRawSocket(inner).register(registry, token, interests), + Self::TCPListener(inner) => inner.deregister(registry), + Self::TCPStream(inner) => SourceFd(inner).deregister(registry), + Self::UDPSocket(inner) => SourceFd(inner).deregister(registry), } } } diff --git a/src/udp.rs b/src/udp.rs index 5de0eed..f22c466 100644 --- a/src/udp.rs +++ b/src/udp.rs @@ -3,7 +3,15 @@ use std::os::fd::{AsRawFd, FromRawFd}; use mio::net::UdpSocket; use pyo3::{IntoPyObject, prelude::*, types::PyBytes}; -use std::{borrow::Cow, cell::RefCell, collections::HashMap, io::ErrorKind, net::SocketAddr, sync::atomic}; +use std::{ + borrow::Cow, + cell::RefCell, + collections::HashMap, + io::ErrorKind, + net::{IpAddr, SocketAddr}, + str::FromStr, + sync::atomic, +}; use crate::{ event_loop::{EventLoop, EventLoopRunState}, @@ -72,13 +80,14 @@ impl UDPTransport { pyloop: &Py, pysock: (i32, i32), proto_factory: PyObject, - remote_addr: Option, + remote_addr_tup: Option<(String, u16)>, ) -> Self { let sock = unsafe { socket2::Socket::from_raw_fd(pysock.0) }; _ = sock.set_nonblocking(true); - let std_socket: std::net::UdpSocket = sock.into(); - let socket = UdpSocket::from_std(std_socket); + let stds: std::net::UdpSocket = sock.into(); + let socket = UdpSocket::from_std(stds); + let remote_addr = remote_addr_tup.map(|v| SocketAddr::new(IpAddr::from_str(&v.0).unwrap(), v.1)); let proto = proto_factory.bind(py).call0().unwrap(); Self::new(py, pyloop.clone_ref(py), socket, proto, pysock.1, remote_addr) @@ -92,7 +101,7 @@ impl UDPTransport { Ok(rself.proto.clone_ref(py)) } - #[inline] + #[inline(always)] fn call_conn_lost(&self, py: Python, exc: Option) { _ = self.protom_conn_lost.call1(py, (exc,)); } @@ -141,13 +150,19 @@ impl UDPTransport { return; } - let event_loop = self.pyloop.get(); - event_loop.udp_socket_rem(self.fd); + self.pyloop.get().udp_socket_rem(self.fd); self.call_conn_lost(py, None); } fn abort(&self, py: Python) { - self.close(py); + if self + .closing + .compare_exchange(false, true, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) + .is_ok() + { + self.pyloop.get().udp_socket_rem(self.fd); + } + self.call_conn_lost(py, None); } fn set_protocol(&self, _protocol: PyObject) -> PyResult<()> { @@ -160,36 +175,24 @@ impl UDPTransport { self.proto.clone_ref(py) } - fn sendto(&self, py: Python, data: Cow<[u8]>, addr: Option) -> PyResult<()> { + // TODO: implement buffered write + fn sendto(&self, data: Cow<[u8]>, addr: Option<(String, u16)>) -> PyResult<()> { if self.closing.load(atomic::Ordering::Relaxed) { return Err(pyo3::exceptions::PyRuntimeError::new_err( "Cannot send on closing transport", )); } - let target_addr = match addr { - Some(addr_obj) => { - // Parse the address from Python tuple (host, port) - let addr_tuple: (String, u16) = addr_obj.extract(py)?; - Some( - format!("{}:{}", addr_tuple.0, addr_tuple.1) - .parse::() - .map_err(|e| pyo3::exceptions::PyValueError::new_err(format!("Invalid address: {e}")))?, - ) - } - None => { - // Get remote addr without holding the borrow longer than necessary - self.state.borrow().remote_addr - } - }; - - match target_addr { + match addr + .map(|v| SocketAddr::new(IpAddr::from_str(&v.0).unwrap(), v.1)) + .or_else(|| self.state.borrow().remote_addr) + { Some(addr) => { // Temporarily borrow just for the send operation match self.state.borrow().socket.send_to(&data, addr) { Ok(_) => Ok(()), Err(err) if err.kind() == ErrorKind::WouldBlock => { - // For UDP, we don't buffer writes like TCP - just drop the packet or return error + // FIXME: For UDP, we don't buffer writes like TCP - just drop the packet or return error Err(pyo3::exceptions::PyBlockingIOError::new_err("Socket would block")) } Err(err) => Err(pyo3::exceptions::PyOSError::new_err(err.to_string())), @@ -200,36 +203,30 @@ impl UDPTransport { } } -pub(crate) struct UDPHandle { +pub(crate) struct UDPReadHandle { fd: usize, } -impl UDPHandle { +impl UDPReadHandle { pub(crate) fn new(fd: usize) -> Self { Self { fd } } } -impl Handle for UDPHandle { - fn run(&self, py: Python, event_loop: &EventLoop, _state: &mut EventLoopRunState) { +impl Handle for UDPReadHandle { + fn run(&self, py: Python, event_loop: &EventLoop, state: &mut EventLoopRunState) { let pytransport = event_loop.get_udp_transport(self.fd, py); let transport = pytransport.borrow(py); - // Read datagrams from the socket - let mut buf = vec![0u8; 65536].into_boxed_slice(); // Max UDP packet size - loop { - // Limit the scope of the borrow - let recv_result = { - let state = transport.state.borrow(); - state.socket.recv_from(&mut buf) - }; - - match recv_result { + match { + let trxstate = transport.state.borrow(); + trxstate.socket.recv_from(&mut state.read_buf) + } { Ok((size, addr)) => { // Call the protocol's datagram_received method // Now state is not borrowed, so sendto can work - transport.call_datagram_received(py, &buf[..size], addr); + transport.call_datagram_received(py, &state.read_buf[..size], addr); } Err(err) if err.kind() == ErrorKind::WouldBlock => { // No more data available From fe14bb51363318084580e814ce1c8d969a2a0b9e Mon Sep 17 00:00:00 2001 From: Giovanni Barillari Date: Wed, 3 Sep 2025 16:44:58 +0200 Subject: [PATCH 4/5] Code lint --- Makefile | 1 + rloop/loop.py | 4 ++-- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/Makefile b/Makefile index b9efcc5..825be12 100644 --- a/Makefile +++ b/Makefile @@ -27,6 +27,7 @@ lint-rust: -D warnings \ -W clippy::pedantic \ -W clippy::dbg_macro \ + -A clippy::blocks_in_conditions \ -A clippy::cast-possible-truncation \ -A clippy::cast-sign-loss \ -A clippy::declare-interior-mutable-const \ diff --git a/rloop/loop.py b/rloop/loop.py index e584057..78479a9 100644 --- a/rloop/loop.py +++ b/rloop/loop.py @@ -693,9 +693,9 @@ async def create_datagram_endpoint( if not infos: raise OSError('getaddrinfo() returned empty list') if local_addr is not None: - addr_info = (infos[0], infos[2], infos[4], None) + addr_info = (infos[0][0], infos[0][2], infos[0][4], None) if remote_addr is not None: - addr_info = (infos[0], infos[2], None, infos[4]) + addr_info = (infos[0][0], infos[0][2], None, infos[0][4]) if not addr_info: raise ValueError('can not get address information') From b0add1c4476c94a9bf90f0bf73244c34736908b7 Mon Sep 17 00:00:00 2001 From: Giovanni Barillari Date: Wed, 10 Sep 2025 12:47:56 +0200 Subject: [PATCH 5/5] Add UDP write buffering --- src/event_loop.rs | 90 ++++++++++---- src/io.rs | 22 ++-- src/udp.rs | 298 +++++++++++++++++++++++++++++++++++++++++----- 3 files changed, 346 insertions(+), 64 deletions(-) diff --git a/src/event_loop.rs b/src/event_loop.rs index 3455698..5471751 100644 --- a/src/event_loop.rs +++ b/src/event_loop.rs @@ -8,8 +8,6 @@ use std::{ }; use anyhow::Result; -#[cfg(unix)] -use mio::unix::SourceFd; use mio::{Interest, Poll, Token, Waker, event, net::TcpListener}; use pyo3::prelude::*; @@ -21,7 +19,7 @@ use crate::{ server::Server, tcp::{TCPReadHandle, TCPServer, TCPServerRef, TCPTransport, TCPWriteHandle}, time::Timer, - udp::{UDPReadHandle, UDPTransport}, + udp::{UDPReadHandle, UDPTransport, UDPWriteHandle}, }; enum IOHandle { @@ -29,7 +27,7 @@ enum IOHandle { Signals, TCPListener(TCPListenerHandleData), TCPStream(Interest), - UDPSocket, + UDPSocket(Interest), } struct PyHandleData { @@ -154,7 +152,7 @@ impl EventLoop { IOHandle::Py(handle) => self.handle_io_py(py, event, handle, &mut cb_handles), IOHandle::TCPListener(handle) => self.handle_io_tcpl(py, handle, &io_handles, &mut cb_handles), IOHandle::TCPStream(_) => self.handle_io_tcps(event, &mut cb_handles), - IOHandle::UDPSocket => self.handle_io_udp(event, &mut cb_handles), + IOHandle::UDPSocket(_) => self.handle_io_udp(event, &mut cb_handles), IOHandle::Signals => self.handle_io_signals(py, &mut state.buf, &mut cb_handles), } } @@ -237,7 +235,7 @@ impl EventLoop { let fd = stream.as_raw_fd() as usize; let token = Token(fd); #[allow(clippy::cast_possible_wrap)] - let mut source = Source::TCPStream(fd as i32); + let mut source = Source::FD(fd as i32); let (pytransport, stream_handle) = handle.server.new_stream(py, stream); transports.insert(fd, pytransport); lstreams.insert(fd); @@ -264,7 +262,9 @@ impl EventLoop { fn handle_io_udp(&self, event: &event::Event, handles_ready: &mut VecDeque) { let fd = event.token().0; if event.is_readable() { - handles_ready.push_back(Box::new(UDPReadHandle::new(fd))); + handles_ready.push_back(Box::new(UDPReadHandle { fd })); + } else if event.is_writable() { + handles_ready.push_back(Box::new(UDPWriteHandle { fd })); } } @@ -352,7 +352,7 @@ impl EventLoop { }, || { #[allow(clippy::cast_possible_wrap)] - let mut source = Source::TCPStream(fd as i32); + let mut source = Source::FD(fd as i32); { let guard_poll = self.io.lock().unwrap(); _ = guard_poll.registry().register(&mut source, token, interest); @@ -418,28 +418,74 @@ impl EventLoop { } #[inline] - pub(crate) fn udp_socket_add(&self, fd: usize) { + pub(crate) fn udp_socket_add(&self, fd: usize, interest: Interest) { let token = Token(fd); - #[allow(clippy::cast_possible_wrap)] - let mut source = Source::UDPSocket(fd as i32); - let guard_poll = self.io.lock().unwrap(); - let _ = guard_poll.registry().register(&mut source, token, Interest::READABLE); - self.handles_io.pin().insert(token, IOHandle::UDPSocket); + self.handles_io.pin().update_or_insert_with( + token, + |io_handle| { + if let IOHandle::UDPSocket(interest_prev) = io_handle { + if *interest_prev == interest { + return IOHandle::UDPSocket(interest); + } + + let interests = *interest_prev | interest; + { + #[allow(clippy::cast_possible_wrap)] + let mut source = Source::FD(fd as i32); + let guard_poll = self.io.lock().unwrap(); + _ = guard_poll.registry().reregister(&mut source, token, interests); + } + return IOHandle::UDPSocket(interests); + } + unreachable!() + }, + || { + #[allow(clippy::cast_possible_wrap)] + let mut source = Source::FD(fd as i32); + { + let guard_poll = self.io.lock().unwrap(); + _ = guard_poll.registry().register(&mut source, token, interest); + } + IOHandle::UDPSocket(interest) + }, + ); } #[inline] - pub(crate) fn udp_socket_rem(&self, fd: usize) { + pub(crate) fn udp_socket_rem(&self, fd: usize, interest: Interest) { let token = Token(fd); - if self.handles_io.pin().remove(&token).is_some() { - let io = self.io.lock().unwrap(); - #[allow(clippy::cast_possible_wrap)] - let _ = io.registry().deregister(&mut SourceFd(&(fd as i32))); + + match self.handles_io.pin().remove_if(&token, |_, io_handle| { + if let IOHandle::UDPSocket(interest_ex) = io_handle { + return *interest_ex == interest; + } + false + }) { + Ok(None) => {} + Ok(_) => { + #[allow(clippy::cast_possible_wrap)] + let mut source = Source::FD(fd as i32); + let guard_poll = self.io.lock().unwrap(); + _ = guard_poll.registry().deregister(&mut source); + } + _ => { + self.handles_io.pin().update(token, |io_handle| { + if let IOHandle::UDPSocket(interest_ex) = io_handle { + let interest_new = interest_ex.remove(interest).unwrap(); + #[allow(clippy::cast_possible_wrap)] + let mut source = Source::FD(fd as i32); + let guard_poll = self.io.lock().unwrap(); + _ = guard_poll.registry().reregister(&mut source, token, interest_new); + return IOHandle::UDPSocket(interest_new); + } + unreachable!() + }); + } } } #[inline] - pub(crate) fn udp_socket_close(&self, _py: Python, fd: usize) { - self.udp_socket_rem(fd); + pub(crate) fn udp_socket_close(&self, fd: usize) { self.udp_transports.pin().remove(&fd); } @@ -1169,7 +1215,7 @@ impl EventLoop { let pytransport = Py::new(py, transport)?; let proto = UDPTransport::attach(&pytransport, py)?; rself.udp_transports.pin().insert(fd, pytransport.clone_ref(py)); - rself.udp_socket_add(fd); + rself.udp_socket_add(fd, Interest::READABLE); Ok((pytransport, proto)) } diff --git a/src/io.rs b/src/io.rs index 1329975..eaf510b 100644 --- a/src/io.rs +++ b/src/io.rs @@ -13,14 +13,14 @@ pub(crate) enum Source { #[cfg(windows)] FD(RawSocket), TCPListener(TcpListener), - #[cfg(unix)] - TCPStream(RawFd), - #[cfg(windows)] - TCPStream(RawSocket), - #[cfg(unix)] - UDPSocket(RawFd), - #[cfg(windows)] - UDPSocket(RawSocket), + // #[cfg(unix)] + // TCPStream(RawFd), + // #[cfg(windows)] + // TCPStream(RawSocket), + // #[cfg(unix)] + // UDPSocket(RawFd), + // #[cfg(windows)] + // UDPSocket(RawSocket), } #[cfg(windows)] @@ -52,8 +52,6 @@ impl MioSource for Source { #[cfg(windows)] Self::FD(inner) => SourceRawSocket(inner).register(registry, token, interests), Self::TCPListener(inner) => inner.register(registry, token, interests), - Self::TCPStream(inner) => SourceFd(inner).register(registry, token, interests), - Self::UDPSocket(inner) => SourceFd(inner).register(registry, token, interests), } } @@ -65,8 +63,6 @@ impl MioSource for Source { #[cfg(windows)] Self::FD(inner) => SourceRawSocket(inner).register(registry, token, interests), Self::TCPListener(inner) => inner.reregister(registry, token, interests), - Self::TCPStream(inner) => SourceFd(inner).reregister(registry, token, interests), - Self::UDPSocket(inner) => SourceFd(inner).reregister(registry, token, interests), } } @@ -78,8 +74,6 @@ impl MioSource for Source { #[cfg(windows)] Self::FD(inner) => SourceRawSocket(inner).register(registry, token, interests), Self::TCPListener(inner) => inner.deregister(registry), - Self::TCPStream(inner) => SourceFd(inner).deregister(registry), - Self::UDPSocket(inner) => SourceFd(inner).deregister(registry), } } } diff --git a/src/udp.rs b/src/udp.rs index f22c466..8fc8ece 100644 --- a/src/udp.rs +++ b/src/udp.rs @@ -1,12 +1,12 @@ #[cfg(unix)] use std::os::fd::{AsRawFd, FromRawFd}; -use mio::net::UdpSocket; -use pyo3::{IntoPyObject, prelude::*, types::PyBytes}; +use mio::{Interest, net::UdpSocket}; +use pyo3::{IntoPyObject, IntoPyObjectExt, prelude::*, types::PyBytes}; use std::{ borrow::Cow, cell::RefCell, - collections::HashMap, + collections::{HashMap, VecDeque}, io::ErrorKind, net::{IpAddr, SocketAddr}, str::FromStr, @@ -16,12 +16,15 @@ use std::{ use crate::{ event_loop::{EventLoop, EventLoopRunState}, handles::Handle, + log::LogExc, sock::SocketWrapper, }; struct UDPTransportState { socket: UdpSocket, remote_addr: Option, + write_buf: VecDeque<(Box<[u8]>, SocketAddr)>, + write_buf_dsize: usize, } #[pyclass(frozen, unsendable, module = "rloop._rloop")] @@ -31,8 +34,11 @@ pub(crate) struct UDPTransport { pyloop: Py, // atomics closing: atomic::AtomicBool, + water_hi: atomic::AtomicUsize, + water_lo: atomic::AtomicUsize, // py protocol fields proto: PyObject, + proto_paused: atomic::AtomicBool, protom_conn_lost: PyObject, protom_datagram_received: PyObject, protom_error_received: PyObject, @@ -51,7 +57,15 @@ impl UDPTransport { remote_addr: Option, ) -> Self { let fd = socket.as_raw_fd() as usize; - let state = UDPTransportState { socket, remote_addr }; + let state = UDPTransportState { + socket, + remote_addr, + write_buf: VecDeque::new(), + write_buf_dsize: 0, + }; + + let wh = 1024 * 64; + let wl = wh / 4; let protom_conn_lost = pyproto.getattr(pyo3::intern!(py, "connection_lost")).unwrap().unbind(); let protom_datagram_received = pyproto @@ -66,7 +80,10 @@ impl UDPTransport { state: RefCell::new(state), pyloop, closing: false.into(), + water_hi: wh.into(), + water_lo: wl.into(), proto, + proto_paused: false.into(), protom_conn_lost, protom_datagram_received, protom_error_received, @@ -101,6 +118,38 @@ impl UDPTransport { Ok(rself.proto.clone_ref(py)) } + #[inline] + fn write_buf_size_decr(pyself: &Py, py: Python) { + let rself = pyself.borrow(py); + if rself.state.borrow().write_buf_dsize <= rself.water_lo.load(atomic::Ordering::Relaxed) + && rself + .proto_paused + .compare_exchange(true, false, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) + .is_ok() + { + Self::proto_resume(pyself, py); + } + } + + #[inline] + fn close_from_write_handle(&self, py: Python, errored: bool) -> bool { + if self.closing.load(atomic::Ordering::Relaxed) { + _ = self.protom_conn_lost.call1( + py, + #[allow(clippy::obfuscated_if_else)] + (errored + .then(|| { + pyo3::exceptions::PyRuntimeError::new_err("socket transport failed") + .into_py_any(py) + .unwrap() + }) + .unwrap_or_else(|| py.None()),), + ); + return true; + } + false + } + #[inline(always)] fn call_conn_lost(&self, py: Python, exc: Option) { _ = self.protom_conn_lost.call1(py, (exc,)); @@ -117,6 +166,90 @@ impl UDPTransport { fn call_error_received(&self, py: Python, exc: PyErr) { _ = self.protom_error_received.call1(py, (exc,)); } + + fn write(pyself: &Py, py: Python, data: &[u8], addr: SocketAddr) { + let rself = pyself.borrow(py); + let mut state = rself.state.borrow_mut(); + + let buf_added = match state.write_buf_dsize { + 0 => { + match rself.state.borrow().socket.send_to(data, addr) { + Ok(written) if written == data.len() => 0, + Ok(written) => { + state.write_buf.push_back(((&data[written..]).into(), addr)); + data.len() - written + } + Err(err) + if err.kind() == std::io::ErrorKind::Interrupted + || err.kind() == std::io::ErrorKind::WouldBlock => + { + state.write_buf.push_back((data.into(), addr)); + data.len() + } + Err(err) => { + if state.write_buf_dsize > 0 { + // reset buf_dsize? + rself.pyloop.get().udp_socket_rem(rself.fd, Interest::WRITABLE); + } + if rself + .closing + .compare_exchange(false, true, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) + .is_ok() + { + rself.pyloop.get().udp_socket_rem(rself.fd, Interest::READABLE); + } + rself.call_conn_lost(py, Some(pyo3::exceptions::PyOSError::new_err(err.to_string()))); + 0 + } + } + } + _ => { + state.write_buf.push_back((data.into(), addr)); + data.len() + } + }; + + if buf_added > 0 { + if state.write_buf_dsize == 0 { + rself.pyloop.get().udp_socket_add(rself.fd, Interest::WRITABLE); + } + state.write_buf_dsize += buf_added; + if state.write_buf_dsize > rself.water_hi.load(atomic::Ordering::Relaxed) + && rself + .proto_paused + .compare_exchange(false, true, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) + .is_ok() + { + Self::proto_pause(pyself, py); + } + } + } + + fn proto_pause(pyself: &Py, py: Python) { + let rself = pyself.borrow(py); + if let Err(err) = rself.proto.call_method0(py, pyo3::intern!(py, "pause_writing")) { + let err_ctx = LogExc::transport( + err, + "protocol.pause_writing() failed".into(), + rself.proto.clone_ref(py), + pyself.clone_ref(py).into_any(), + ); + _ = rself.pyloop.get().log_exception(py, err_ctx); + } + } + + fn proto_resume(pyself: &Py, py: Python) { + let rself = pyself.borrow(py); + if let Err(err) = rself.proto.call_method0(py, pyo3::intern!(py, "resume_writing")) { + let err_ctx = LogExc::transport( + err, + "protocol.resume_writing() failed".into(), + rself.proto.clone_ref(py), + pyself.clone_ref(py).into_any(), + ); + _ = rself.pyloop.get().log_exception(py, err_ctx); + } + } } #[pymethods] @@ -150,17 +283,24 @@ impl UDPTransport { return; } - self.pyloop.get().udp_socket_rem(self.fd); - self.call_conn_lost(py, None); + let event_loop = self.pyloop.get(); + event_loop.udp_socket_rem(self.fd, Interest::READABLE); + if self.state.borrow().write_buf_dsize == 0 { + event_loop.udp_socket_rem(self.fd, Interest::WRITABLE); + self.call_conn_lost(py, None); + } } fn abort(&self, py: Python) { + if self.state.borrow().write_buf_dsize > 0 { + self.pyloop.get().udp_socket_rem(self.fd, Interest::WRITABLE); + } if self .closing .compare_exchange(false, true, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) .is_ok() { - self.pyloop.get().udp_socket_rem(self.fd); + self.pyloop.get().udp_socket_rem(self.fd, Interest::READABLE); } self.call_conn_lost(py, None); } @@ -175,28 +315,72 @@ impl UDPTransport { self.proto.clone_ref(py) } - // TODO: implement buffered write - fn sendto(&self, data: Cow<[u8]>, addr: Option<(String, u16)>) -> PyResult<()> { - if self.closing.load(atomic::Ordering::Relaxed) { + #[pyo3(signature = (high = None, low = None))] + fn set_write_buffer_limits(pyself: Py, py: Python, high: Option, low: Option) -> PyResult<()> { + let wh = match high { + None => match low { + None => 1024 * 64, + Some(v) => v * 4, + }, + Some(v) => v, + }; + let wl = match low { + None => wh / 4, + Some(v) => v, + }; + + if wh < wl { + return Err(pyo3::exceptions::PyValueError::new_err( + "high must be >= low must be >= 0", + )); + } + + let rself = pyself.borrow(py); + rself.water_hi.store(wh, atomic::Ordering::Relaxed); + rself.water_lo.store(wl, atomic::Ordering::Relaxed); + + if rself.state.borrow().write_buf_dsize > wh + && rself + .proto_paused + .compare_exchange(false, true, atomic::Ordering::Relaxed, atomic::Ordering::Relaxed) + .is_ok() + { + Self::proto_pause(&pyself, py); + } + + Ok(()) + } + + fn get_write_buffer_size(&self) -> usize { + self.state.borrow().write_buf_dsize + } + + fn get_write_buffer_limits(&self) -> (usize, usize) { + ( + self.water_lo.load(atomic::Ordering::Relaxed), + self.water_hi.load(atomic::Ordering::Relaxed), + ) + } + + fn sendto(pyself: Py, py: Python, data: Cow<[u8]>, addr: Option<(String, u16)>) -> PyResult<()> { + let rself = pyself.borrow(py); + + if rself.closing.load(atomic::Ordering::Relaxed) { return Err(pyo3::exceptions::PyRuntimeError::new_err( "Cannot send on closing transport", )); } + if data.is_empty() { + return Ok(()); + } match addr .map(|v| SocketAddr::new(IpAddr::from_str(&v.0).unwrap(), v.1)) - .or_else(|| self.state.borrow().remote_addr) + .or_else(|| rself.state.borrow().remote_addr) { Some(addr) => { - // Temporarily borrow just for the send operation - match self.state.borrow().socket.send_to(&data, addr) { - Ok(_) => Ok(()), - Err(err) if err.kind() == ErrorKind::WouldBlock => { - // FIXME: For UDP, we don't buffer writes like TCP - just drop the packet or return error - Err(pyo3::exceptions::PyBlockingIOError::new_err("Socket would block")) - } - Err(err) => Err(pyo3::exceptions::PyOSError::new_err(err.to_string())), - } + Self::write(&pyself, py, &data, addr); + Ok(()) } None => Err(pyo3::exceptions::PyValueError::new_err("No remote address specified")), } @@ -204,13 +388,7 @@ impl UDPTransport { } pub(crate) struct UDPReadHandle { - fd: usize, -} - -impl UDPReadHandle { - pub(crate) fn new(fd: usize) -> Self { - Self { fd } - } + pub fd: usize, } impl Handle for UDPReadHandle { @@ -239,10 +417,74 @@ impl Handle for UDPReadHandle { // Other error - call error_received and close let py_err = pyo3::exceptions::PyOSError::new_err(err.to_string()); transport.call_error_received(py, py_err); - event_loop.udp_socket_close(py, self.fd); + event_loop.udp_socket_close(self.fd); + break; + } + } + } + } +} + +pub(crate) struct UDPWriteHandle { + pub fd: usize, +} + +impl UDPWriteHandle { + #[inline] + fn write(&self, transport: &UDPTransport) -> Option { + let mut ret = 0; + let mut state = transport.state.borrow_mut(); + + while let Some((data, addr)) = state.write_buf.pop_front() { + match state.socket.send_to(&data, addr) { + Ok(written) if written < data.len() => { + state.write_buf.push_front(((&data[written..]).into(), addr)); + ret += written; break; } + Ok(written) => ret += written, + Err(err) if err.kind() == std::io::ErrorKind::Interrupted => { + state.write_buf.push_front((data, addr)); + } + Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { + state.write_buf.push_front((data, addr)); + break; + } + _ => { + state.write_buf.clear(); + state.write_buf_dsize = 0; + return None; + } } } + state.write_buf_dsize -= ret; + Some(ret) + } +} + +impl Handle for UDPWriteHandle { + fn run(&self, py: Python, event_loop: &EventLoop, _state: &mut EventLoopRunState) { + let pytransport = event_loop.get_udp_transport(self.fd, py); + let transport = pytransport.borrow(py); + let stream_close; + + if let Some(written) = self.write(&transport) { + if written > 0 { + UDPTransport::write_buf_size_decr(&pytransport, py); + } + stream_close = match transport.state.borrow().write_buf.is_empty() { + true => transport.close_from_write_handle(py, false), + false => false, + }; + } else { + stream_close = transport.close_from_write_handle(py, true); + } + + if transport.state.borrow().write_buf.is_empty() { + event_loop.udp_socket_rem(self.fd, Interest::WRITABLE); + } + if stream_close { + event_loop.udp_socket_close(self.fd); + } } }