diff --git a/README.md b/README.md index 22045ac..125b435 100644 --- a/README.md +++ b/README.md @@ -223,6 +223,35 @@ If you do not install an error handler, the task still fails cleanly and the runtime keeps its internal state consistent, but adding `setErrorHandler(...)` is the recommended way to make these failures visible in applications. +### Waking a Blocked Scheduler from Another Thread + +Kernels may provide an opaque wakeup channel for code that must request work +such as server shutdown while the scheduler is blocked in I/O readiness: + +```python +kernel = Unix() +if not kernel.supports_wakeup_channel(): + raise RuntimeError("cross-thread scheduler wakeup is unavailable") + +wakeup = kernel.create_wakeup_channel() + +async def watch_shutdown(task): + await task.wait_readable(wakeup.wait_object) + wakeup.drain() + # Apply the application-owned shutdown request on the scheduler thread. +``` + +Call `wakeup.notify()` from the external thread. Notifications are nonblocking +and coalesce until the scheduler calls `drain()`. The owner must call `close()` +after its scheduler wait has been detached; repeated close, notify, and drain +calls during teardown are safe. + +`Unix` supports this contract when its socket module provides a callable +`socketpair()`. Generic `MicroPythonKernel` deliberately reports it unsupported: +a polling backend or socket-pair-shaped attribute alone does not establish safe +cross-thread behavior on a constrained port. TCP serving and shutdown initiated +by a task already running on the scheduler do not require this capability. + ## Configuration The runtime now uses a first-class config object backed by @@ -460,6 +489,15 @@ The kernel layer is deliberately generic. Protocol clients such as HTTPS, Redis, MQTT, RabbitMQ/AMQP, and Kafka should be built on top of the shared TCP/TLS socket surface rather than requiring protocol-specific kernel methods. +Passive TCP consumers use the kernel boundary as well: check +`supports_tcp_server()` before resolving or opening anything, pass the opaque +record returned by `resolve_passive_address()` unchanged to both `socket_open()` +and `socket_bind()`, then use the kernel's listen, accept, address-inspection, +and close operations. Address reuse has its own capability check because some +MicroPython ports support listeners without exposing `SO_REUSEADDR` constants. +The web app demo shows the complete setup and rollback pattern without importing +platform socket APIs. + ## Clients The current setup now includes first-party smallOS-native helpers for HTTP, diff --git a/SmallPackage/Kernel.py b/SmallPackage/Kernel.py index e01bed6..2a8a2b6 100644 --- a/SmallPackage/Kernel.py +++ b/SmallPackage/Kernel.py @@ -24,9 +24,38 @@ if TYPE_CHECKING: from collections.abc import Iterable, Mapping, Sequence from typing import Any, cast + from ._types import WakeupChannelLike + from ._types import SocketBuffer, SocketOperation, SocketRetryMode _UNSET = object() +_MAX_WAKEUP_IO_ATTEMPTS = 8 +_SOCKET_OPERATIONS = ('accept', 'recv', 'send', 'handshake') + + +def _validate_socket_operation(operation): + """Reject retry classifications that do not name a supported operation.""" + if operation not in _SOCKET_OPERATIONS: + raise ValueError( + 'Unknown socket operation {!r}; expected one of: {}.'.format( + operation, ', '.join(_SOCKET_OPERATIONS) + ) + ) + + +class _PassiveAddressInfo: + """Kernel-owned address record for a passive TCP endpoint.""" + + def __init__(self, owner, address_info): + self.owner = owner + self.address_info = address_info + + +def _unwrap_passive_address(owner, address_info): + """Return backend address data only for a record created by ``owner``.""" + if not isinstance(address_info, _PassiveAddressInfo) or address_info.owner is not owner: + raise ValueError('Passive address record belongs to a different kernel.') + return address_info.address_info def _import_first(*module_names: str) -> Any | None: @@ -369,6 +398,117 @@ def close(self): self._closed = True +class _SocketWakeupChannel: + """Coalescing cross-thread wakeup backed by a non-blocking socket pair.""" + + def __init__(self, reader, writer, lock): + self._reader = reader + self._writer = writer + self._lock = lock + self._pending = False + self._terminal = False + self._closed = False + self._reader_released = False + self._writer_released = False + + @property + def wait_object(self): + """Return the opaque object registered for readable readiness.""" + return self._reader + + def _make_terminal(self): + """Close the writer so EOF wakes the reader, without false latching.""" + try: + self._writer.close() + except BaseException: + # A later notify must retry when neither send nor close established a + # visible wakeup. + return False + self._writer_released = True + self._terminal = True + self._pending = True + return True + + def notify(self): + """Publish one coalescing notification without blocking indefinitely.""" + with self._lock: + if self._closed or self._pending or self._terminal: + return + + for _attempt in range(_MAX_WAKEUP_IO_ATTEMPTS): + try: + sent = self._writer.send(b'\x00') + except InterruptedError: + continue + except BlockingIOError: + # A full non-blocking channel is already readable. + self._pending = True + return + except Exception: + self._make_terminal() + return + + try: + made_progress = sent > 0 + except Exception: + self._make_terminal() + return + if made_progress: + self._pending = True + return + self._make_terminal() + return + + # Repeated interruption is terminal only if closing the writer really + # succeeds and therefore makes EOF observable on the reader. + self._make_terminal() + + def drain(self): + """Clear a normal pending byte while bounding interrupted receive retries.""" + with self._lock: + if self._closed: + return + + for _attempt in range(_MAX_WAKEUP_IO_ATTEMPTS): + try: + data = self._reader.recv(4096) + except InterruptedError: + continue + except BlockingIOError: + if not self._terminal: + self._pending = False + return + except Exception: + # Preserve pending state because the notification may remain unread. + return + + if not data: + self._terminal = True + self._pending = True + return + + # The bounded loop intentionally preserves pending state. A caller may + # retry drain, but notify cannot race in and create a lost wakeup. + + def close(self): + """Idempotently release both endpoints, retrying prior close failures.""" + with self._lock: + self._closed = True + self._pending = False + self._terminal = True + for endpoint, released_name in ( + (self._reader, '_reader_released'), + (self._writer, '_writer_released'), + ): + if getattr(self, released_name): + continue + try: + endpoint.close() + except BaseException: + continue + setattr(self, released_name, True) + + def detect_micropython_machine_name(sys_mod: Any = None, os_mod: Any = None) -> str: """ Best-effort lookup of the active board/firmware machine name. @@ -512,6 +652,22 @@ def supports_external_wait_objects(self) -> bool: """Whether ``io_wait`` can wake on adapter-owned readiness objects.""" return False + def supports_wakeup_channel(self) -> bool: + """Whether this kernel can wake a blocked scheduler from another thread.""" + return False + + def create_wakeup_channel(self) -> WakeupChannelLike: + """Create an opaque cross-thread scheduler wakeup channel.""" + raise NotImplementedError('Cross-thread wakeup channels are not supported.') + + def supports_tcp_server(self) -> bool: + """Whether this kernel implements the complete passive TCP contract.""" + return False + + def supports_reuse_address(self) -> bool: + """Whether this kernel can configure address reuse on a listener.""" + return False + def validate_io_wait_object( self, obj: Any ) -> tuple[bool, BaseException | None]: @@ -544,11 +700,32 @@ def validate_io_wait_object( def resolve_address(self, host: str, port: int) -> Any: return None + def resolve_passive_address(self, host: str, port: int) -> Any: + raise NotImplementedError('Passive TCP address resolution is not supported.') + def socket_open(self, address_info: Any) -> Any: - return None + raise NotImplementedError('TCP sockets are not supported.') def socket_setblocking(self, sock: Any, flag: bool) -> None: - return + raise NotImplementedError('TCP sockets are not supported.') + + def socket_set_reuse_address(self, sock: Any, enabled: bool) -> None: + raise NotImplementedError('TCP address reuse is not supported.') + + def socket_bind(self, sock: Any, address_info: Any) -> None: + raise NotImplementedError('TCP listeners are not supported.') + + def socket_listen(self, sock: Any, backlog: int) -> None: + raise NotImplementedError('TCP listeners are not supported.') + + def socket_accept(self, listener: Any) -> tuple[Any, Any]: + raise NotImplementedError('TCP listeners are not supported.') + + def socket_local_address(self, sock: Any) -> Any: + raise NotImplementedError('Socket address inspection is not supported.') + + def socket_peer_address(self, sock: Any) -> Any | None: + raise NotImplementedError('Socket address inspection is not supported.') def socket_connect(self, sock: Any, sockaddr: Any) -> bool: return True @@ -556,14 +733,15 @@ def socket_connect(self, sock: Any, sockaddr: Any) -> bool: def socket_connection_error(self, sock: Any) -> int: return 0 - def socket_send(self, sock: Any, data: bytes) -> int: + def socket_send(self, sock: Any, data: SocketBuffer) -> int: + """Send a bytes-like buffer without requiring an intermediate copy.""" return 0 def socket_recv(self, sock: Any, buffer_size: int) -> bytes: return b'' def socket_close(self, sock: Any) -> None: - return + raise NotImplementedError('TCP sockets are not supported.') def socket_wrap_tls_client( self, @@ -593,6 +771,17 @@ def socket_needs_read(self, exc: BaseException) -> bool: def socket_needs_write(self, exc: BaseException) -> bool: return False + def socket_retry_mode( + self, exc: BaseException, operation: SocketOperation + ) -> SocketRetryMode: + """Return the readiness direction needed to retry one socket operation.""" + _validate_socket_operation(operation) + if self.socket_needs_write(exc): + return 'write' + if self.socket_needs_read(exc): + return 'write' if operation == 'send' else 'read' + return None + def _poll_lookup_key(self, obj: Any) -> Any: """ Return the identity key used to map poll events back to registered objects. @@ -623,6 +812,7 @@ def __init__(self): import socket import ssl import sys + import threading import time self._errno = errno @@ -640,6 +830,7 @@ def __init__(self): self._socket = socket self._ssl = ssl self._sys = sys + self._lock_factory = threading.Lock self._time = time self._poll_factory = getattr(select, 'poll', None) return @@ -736,10 +927,95 @@ def validate_io_wait_object(self, obj: Any) -> tuple[bool, BaseException | None] def supports_external_wait_objects(self) -> bool: return True + def supports_wakeup_channel(self) -> bool: + """Require a callable primitive rather than mere attribute presence.""" + return callable(getattr(self._socket, 'socketpair', None)) + + def create_wakeup_channel(self) -> WakeupChannelLike: + """Create a non-blocking socket-pair channel owned by this kernel.""" + socket_pair = getattr(self._socket, 'socketpair', None) + if not callable(socket_pair): + raise NotImplementedError( + 'This Unix platform does not provide a callable socketpair().' + ) + + acquired = [] + try: + pair = socket_pair() + if TYPE_CHECKING: + pair = cast("Any", pair) + iterator = iter(pair) + reader = next(iterator) + acquired.append(reader) + writer = next(iterator) + acquired.append(writer) + try: + extra = next(iterator) + except StopIteration: + pass + else: + acquired.append(extra) + raise ValueError('socketpair() must return exactly two endpoints.') + if reader is writer: + raise ValueError('socketpair() endpoints must be distinct objects.') + for endpoint, required in ( + (reader, ('setblocking', 'recv', 'close')), + (writer, ('setblocking', 'send', 'close')), + ): + missing = [ + name for name in required + if not callable(getattr(endpoint, name, None)) + ] + if missing: + raise TypeError( + 'socketpair() endpoint is missing callable operations: {}.'.format( + ', '.join(missing) + ) + ) + reader.setblocking(False) + writer.setblocking(False) + lock = self._lock_factory() + except BaseException: + released = [] + for endpoint in reversed(acquired): + if any(endpoint is released_endpoint for released_endpoint in released): + continue + released.append(endpoint) + try: + closer = getattr(endpoint, 'close', None) + if callable(closer): + closer() + except BaseException: + pass + raise + return _SocketWakeupChannel(reader, writer, lock) + + def supports_tcp_server(self) -> bool: + return True + + def supports_reuse_address(self) -> bool: + return ( + hasattr(self._socket, 'SOL_SOCKET') + and hasattr(self._socket, 'SO_REUSEADDR') + and callable(getattr(self._socket.socket, 'setsockopt', None)) + ) + def resolve_address(self, host, port): return self._socket.getaddrinfo(host, port, type=self._socket.SOCK_STREAM)[0] + def resolve_passive_address(self, host, port): + passive_host = None if host == '' else host + address_info = self._socket.getaddrinfo( + passive_host, + port, + type=self._socket.SOCK_STREAM, + flags=self._socket.AI_PASSIVE, + )[0] + return _PassiveAddressInfo(self, address_info) + def socket_open(self, address_info): + if isinstance(address_info, _PassiveAddressInfo): + address_info = _unwrap_passive_address(self, address_info) family, socktype, proto, _, _ = address_info return self._socket.socket(family, socktype, proto) @@ -747,6 +1023,39 @@ def socket_setblocking(self, sock, flag): sock.setblocking(flag) return + def socket_set_reuse_address(self, sock, enabled): + if not self.supports_reuse_address(): + raise NotImplementedError('SO_REUSEADDR is not supported on this platform.') + sock.setsockopt( + self._socket.SOL_SOCKET, + self._socket.SO_REUSEADDR, + 1 if enabled else 0, + ) + return + + def socket_bind(self, sock, address_info): + address_info = _unwrap_passive_address(self, address_info) + sock.bind(address_info[4]) + return + + def socket_listen(self, sock, backlog): + sock.listen(backlog) + return + + def socket_accept(self, listener): + return listener.accept() + + def socket_local_address(self, sock): + return sock.getsockname() + + def socket_peer_address(self, sock): + try: + return sock.getpeername() + except OSError as exc: + if self._extract_errno(exc) == self._errno.ENOTCONN: + return None + raise + def socket_connect(self, sock, sockaddr): err = sock.connect_ex(sockaddr) pending = { @@ -766,7 +1075,7 @@ def socket_connect(self, sock, sockaddr): def socket_connection_error(self, sock): return sock.getsockopt(self._socket.SOL_SOCKET, self._socket.SO_ERROR) - def socket_send(self, sock, data): + def socket_send(self, sock: Any, data: SocketBuffer) -> int: return sock.send(data) def socket_recv(self, sock, buffer_size): @@ -811,6 +1120,25 @@ def socket_needs_read(self, exc): def socket_needs_write(self, exc): return isinstance(exc, self._ssl.SSLWantWriteError) + def socket_retry_mode( + self, exc: BaseException, operation: SocketOperation + ) -> SocketRetryMode: + _validate_socket_operation(operation) + if isinstance(exc, self._ssl.SSLWantReadError): + return 'read' + if isinstance(exc, self._ssl.SSLWantWriteError): + return 'write' + would_block = { + value for value in ( + getattr(self._errno, 'EAGAIN', None), + getattr(self._errno, 'EWOULDBLOCK', None), + ) if value is not None + } + err = self._extract_errno(exc) + if err in would_block or (isinstance(exc, BlockingIOError) and err is None): + return 'write' if operation == 'send' else 'read' + return None + class MicroPythonKernel(Kernel): ''' @@ -848,6 +1176,7 @@ def __init__( self._ssl = _import_first('ssl', 'ussl') self._sys = modules.get('sys') or _import_first('sys') self._os = modules.get('os') or _import_first('os') + self._errno = modules.get('errno') or _import_first('errno', 'uerrno') self._network = modules.get('network') or _import_first('network') self._rp2 = modules.get('rp2') or _import_first('rp2') self._machine = modules.get('machine') or _import_first('machine') @@ -936,17 +1265,145 @@ def create_io_wait_set(self): def supports_external_wait_objects(self) -> bool: return bool(self._poll_factory) + def supports_wakeup_channel(self) -> bool: + # Poll support or a socketpair-shaped attribute does not prove that a + # constrained port can safely signal it from another thread or interrupt. + return False + + def create_wakeup_channel(self) -> WakeupChannelLike: + raise NotImplementedError( + 'Cross-thread wakeup channels are not supported by MicroPythonKernel.' + ) + + def supports_tcp_server(self) -> bool: + if not callable(getattr(self._socket, 'getaddrinfo', None)): + return False + socket_factory = getattr(self._socket, 'socket', None) + if not callable(socket_factory): + return False + required = ( + 'setblocking', 'bind', 'listen', 'accept', 'getsockname', + 'getpeername', 'send', 'recv', 'close', + ) + return all(callable(getattr(socket_factory, name, None)) for name in required) + + def supports_reuse_address(self) -> bool: + socket_factory = getattr(self._socket, 'socket', None) + return ( + hasattr(self._socket, 'SOL_SOCKET') + and hasattr(self._socket, 'SO_REUSEADDR') + and callable(getattr(socket_factory, 'setsockopt', None)) + ) + def resolve_address(self, host, port): return self._socket.getaddrinfo(host, port)[0] + def resolve_passive_address(self, host, port): + if not self.supports_tcp_server(): + raise NotImplementedError('TCP server support is not available on this port.') + passive_host = '0.0.0.0' if host == '' else host + sock_stream = getattr(self._socket, 'SOCK_STREAM', None) + ai_passive = getattr(self._socket, 'AI_PASSIVE', None) + attempts = [] + if sock_stream is not None and ai_passive is not None: + attempts.append((passive_host, port, 0, sock_stream, 0, ai_passive)) + if sock_stream is not None: + attempts.append((passive_host, port, 0, sock_stream)) + attempts.append((passive_host, port)) + last_error = None + for arguments in attempts: + try: + resolved = self._socket.getaddrinfo(*arguments)[0] + break + except TypeError as exc: + last_error = exc + else: + if last_error is None: + raise RuntimeError('Passive TCP address resolution did not run.') + raise last_error + return _PassiveAddressInfo(self, resolved) + def socket_open(self, address_info): + is_passive = isinstance(address_info, _PassiveAddressInfo) + if is_passive: + address_info = _unwrap_passive_address(self, address_info) family, socktype, proto, _, _ = address_info - return self._socket.socket(family, socktype, proto) + sock = self._socket.socket(family, socktype, proto) + if not is_passive: + return sock + required = ( + 'setblocking', 'bind', 'listen', 'accept', 'getsockname', + 'getpeername', 'send', 'recv', 'close', + ) + missing = [name for name in required if not callable(getattr(sock, name, None))] + if not missing: + return sock + closer = getattr(sock, 'close', None) + if callable(closer): + try: + closer() + except BaseException: + pass + raise NotImplementedError( + 'TCP server socket is missing required operations: {}.'.format(', '.join(missing)) + ) def socket_setblocking(self, sock, flag): sock.setblocking(flag) return + def socket_set_reuse_address(self, sock, enabled): + if not self.supports_reuse_address(): + raise NotImplementedError('SO_REUSEADDR is not supported on this port.') + setter = getattr(sock, 'setsockopt', None) + if not callable(setter): + raise NotImplementedError('Socket instance does not support SO_REUSEADDR.') + setter( + self._socket.SOL_SOCKET, + self._socket.SO_REUSEADDR, + 1 if enabled else 0, + ) + return + + def socket_bind(self, sock, address_info): + address_info = _unwrap_passive_address(self, address_info) + binder = getattr(sock, 'bind', None) + if not callable(binder): + raise NotImplementedError('TCP bind is not supported on this port.') + binder(address_info[4]) + return + + def socket_listen(self, sock, backlog): + listener = getattr(sock, 'listen', None) + if not callable(listener): + raise NotImplementedError('TCP listen is not supported on this port.') + listener(backlog) + return + + def socket_accept(self, listener): + acceptor = getattr(listener, 'accept', None) + if not callable(acceptor): + raise NotImplementedError('TCP accept is not supported on this port.') + return acceptor() + + def socket_local_address(self, sock): + getter = getattr(sock, 'getsockname', None) + if not callable(getter): + raise NotImplementedError('Local address inspection is not supported.') + return getter() + + def socket_peer_address(self, sock): + getter = getattr(sock, 'getpeername', None) + if not callable(getter): + raise NotImplementedError('Peer address inspection is not supported.') + try: + return getter() + except OSError as exc: + not_connected = getattr(self._errno, 'ENOTCONN', None) + if not_connected is not None and self._extract_errno(exc) == not_connected: + return None + raise + def socket_connect(self, sock, sockaddr): try: sock.connect(sockaddr) @@ -957,7 +1414,7 @@ def socket_connect(self, sock, sockaddr): pending = { 115, 11, } - if err in pending or self.socket_needs_read(exc) or self.socket_needs_write(exc): + if err in pending or isinstance(exc, BlockingIOError): return False raise @@ -966,7 +1423,7 @@ def socket_connection_error(self, sock): so_error = getattr(self._socket, 'SO_ERROR', 4) return sock.getsockopt(sol_socket, so_error) - def socket_send(self, sock, data): + def socket_send(self, sock: Any, data: SocketBuffer) -> int: return sock.send(data) def socket_recv(self, sock, buffer_size): @@ -1050,6 +1507,27 @@ def socket_needs_write(self, exc): return True return False + def socket_retry_mode( + self, exc: BaseException, operation: SocketOperation + ) -> SocketRetryMode: + _validate_socket_operation(operation) + want_read_error = getattr(self._ssl, 'SSLWantReadError', None) + if isinstance(want_read_error, type) and isinstance(exc, want_read_error): + return 'read' + want_write_error = getattr(self._ssl, 'SSLWantWriteError', None) + if isinstance(want_write_error, type) and isinstance(exc, want_write_error): + return 'write' + err = self._extract_errno(exc) + pending = {11} + if self._errno is not None: + for name in ('EAGAIN', 'EWOULDBLOCK'): + value = getattr(self._errno, name, None) + if value is not None: + pending.add(value) + if err in pending or (isinstance(exc, BlockingIOError) and err is None): + return 'write' if operation == 'send' else 'read' + return None + def machine_name(self): ''' Return the current firmware/board descriptor when the platform exposes one. diff --git a/SmallPackage/_types.py b/SmallPackage/_types.py index 6589d5e..0c38612 100644 --- a/SmallPackage/_types.py +++ b/SmallPackage/_types.py @@ -7,13 +7,16 @@ from __future__ import annotations from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Protocol, TypeAlias, TypedDict, TypeVar +from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeAlias, TypedDict, TypeVar if TYPE_CHECKING: from .SmallOS import SmallOS T = TypeVar("T") +SocketBuffer: TypeAlias = bytes | bytearray | memoryview +SocketOperation: TypeAlias = Literal["accept", "recv", "send", "handshake"] +SocketRetryMode: TypeAlias = Literal["read", "write"] | None class TaskLike(Protocol): @@ -42,6 +45,50 @@ def io_wait( ) -> tuple[Sequence[Any], Sequence[Any]]: ... +class WakeupChannelLike(Protocol): + """Opaque readiness channel used to wake a scheduler across threads.""" + + @property + def wait_object(self) -> Any: ... + + def notify(self) -> None: ... + def drain(self) -> None: ... + def close(self) -> None: ... + + +class WakeupKernelLike(Protocol): + """Optional kernel boundary for cross-thread scheduler wakeups.""" + + def supports_wakeup_channel(self) -> bool: ... + def create_wakeup_channel(self) -> WakeupChannelLike: ... + + +class PassiveTCPKernelLike(Protocol): + """Platform-neutral passive TCP operations used by server consumers.""" + + def supports_tcp_server(self) -> bool: ... + def supports_reuse_address(self) -> bool: ... + def resolve_passive_address(self, host: str, port: int) -> Any: ... + def socket_open(self, address_info: Any) -> Any: ... + def socket_setblocking(self, sock: Any, flag: bool) -> None: ... + def socket_set_reuse_address(self, sock: Any, enabled: bool) -> None: ... + def socket_bind(self, sock: Any, address_info: Any) -> None: ... + def socket_listen(self, sock: Any, backlog: int) -> None: ... + def socket_accept(self, listener: Any) -> tuple[Any, Any]: ... + def socket_local_address(self, sock: Any) -> Any: ... + def socket_peer_address(self, sock: Any) -> Any | None: ... + def socket_close(self, sock: Any) -> None: ... + + +class SocketKernelLike(Protocol): + """Optional kernel boundary for operation-aware byte-stream I/O.""" + + def socket_send(self, sock: Any, data: SocketBuffer) -> int: ... + def socket_retry_mode( + self, exc: BaseException, operation: SocketOperation + ) -> SocketRetryMode: ... + + class AdapterCompletionLike(Protocol): """Structural completion record consumed by the scheduler.""" diff --git a/SmallPackage/clients/SmallStream.py b/SmallPackage/clients/SmallStream.py index 2cbcd2f..3125f6c 100644 --- a/SmallPackage/clients/SmallStream.py +++ b/SmallPackage/clients/SmallStream.py @@ -100,10 +100,11 @@ async def connect(self): kernel.socket_do_handshake(sock) break except Exception as exc: - if kernel.socket_needs_read(exc): + retry_mode = kernel.socket_retry_mode(exc, "handshake") + if retry_mode == "read": await self.task.wait_readable(sock) continue - if kernel.socket_needs_write(exc): + if retry_mode == "write": await self.task.wait_writable(sock) continue raise @@ -135,7 +136,7 @@ async def send_all(self, data): if not self._connected or self.sock is None: await self.connect() - view = memoryview(bytes(data)) + view = data if isinstance(data, memoryview) else memoryview(data) while view: try: sent = self.kernel.socket_send(self.sock, view) @@ -143,10 +144,11 @@ async def send_all(self, data): raise StreamClosedError("socket closed while sending data") view = view[sent:] except Exception as exc: - if self.kernel.socket_needs_read(exc): + retry_mode = self.kernel.socket_retry_mode(exc, "send") + if retry_mode == "read": await self.task.wait_readable(self.sock) continue - if self.kernel.socket_needs_write(exc): + if retry_mode == "write": await self.task.wait_writable(self.sock) continue raise @@ -164,10 +166,11 @@ async def recv_some(self, size=4096): raise StreamClosedError("socket closed while reading data") return bytes(chunk) except Exception as exc: - if self.kernel.socket_needs_read(exc): + retry_mode = self.kernel.socket_retry_mode(exc, "recv") + if retry_mode == "read": await self.task.wait_readable(self.sock) continue - if self.kernel.socket_needs_write(exc): + if retry_mode == "write": await self.task.wait_writable(self.sock) continue raise diff --git a/demos/web_app_demo.py b/demos/web_app_demo.py index 6244b30..bb1c37a 100644 --- a/demos/web_app_demo.py +++ b/demos/web_app_demo.py @@ -184,16 +184,17 @@ def _home_page(): async def _send_all(task, sock, data): """Write all bytes to a non-blocking socket using smallOS waits.""" kernel = task.OS.kernel - remaining = memoryview(bytes(data)) + remaining = data if isinstance(data, memoryview) else memoryview(data) while remaining: try: sent = kernel.socket_send(sock, remaining) except Exception as exc: - if kernel.socket_needs_read(exc): + retry_mode = kernel.socket_retry_mode(exc, "send") + if retry_mode == "read": await task.wait_readable(sock) continue - if kernel.socket_needs_write(exc): + if retry_mode == "write": await task.wait_writable(sock) continue raise @@ -212,10 +213,11 @@ async def _read_request_head(task, sock): try: chunk = kernel.socket_recv(sock, 1024) except Exception as exc: - if kernel.socket_needs_read(exc): + retry_mode = kernel.socket_retry_mode(exc, "recv") + if retry_mode == "read": await task.wait_readable(sock) continue - if kernel.socket_needs_write(exc): + if retry_mode == "write": await task.wait_writable(sock) continue raise @@ -340,53 +342,78 @@ async def metrics_task(task, state): await task.sleep(METRICS_INTERVAL_SECONDS) +def _close_after_setup_failure(kernel, stream): + """Best-effort rollback that never replaces the setup exception.""" + try: + kernel.socket_close(stream) + except BaseException: + pass + + +def _open_listener(kernel, host, port, backlog): + """Construct one listener entirely through the passive TCP kernel contract.""" + if not kernel.supports_tcp_server(): + raise NotImplementedError("This kernel does not support passive TCP servers.") + + address_info = kernel.resolve_passive_address(host, port) + listener = kernel.socket_open(address_info) + try: + if kernel.supports_reuse_address(): + kernel.socket_set_reuse_address(listener, True) + kernel.socket_bind(listener, address_info) + kernel.socket_listen(listener, backlog) + kernel.socket_setblocking(listener, False) + except BaseException: + _close_after_setup_failure(kernel, listener) + raise + return listener + + +def _dispatch_client(task, client_sock, client_addr, state): + """Transfer one accepted stream to a handler or roll it back exactly once.""" + kernel = task.OS.kernel + try: + kernel.socket_setblocking(client_sock, False) + task.spawn( + web_client_handler, + priority=max(1, task.priority - 1), + name="http_client", + args=(client_sock, client_addr, state), + ) + except BaseException: + _close_after_setup_failure(kernel, client_sock) + raise + + async def web_server_task(task, state): """Run a small cooperative HTTP server on one non-blocking listener.""" kernel = task.OS.kernel if kernel is None: raise RuntimeError("web_server_task requires a kernel-enabled runtime.") - address_info = kernel.resolve_address(HOST, PORT) - listener = kernel.socket_open(address_info) - - # Reuse address when available so restarting the demo is less annoying. - if hasattr(listener, "setsockopt") and hasattr(listener, "SOL_SOCKET") and hasattr(listener, "SO_REUSEADDR"): - try: - listener.setsockopt(listener.SOL_SOCKET, listener.SO_REUSEADDR, 1) - except Exception: - pass - - listener.bind(address_info[4]) - listener.listen(LISTEN_BACKLOG) - kernel.socket_setblocking(listener, False) - - task.OS.print("smallOS web app running on http://{}:{}/\n".format(HOST, PORT)) - task.OS.print("routes: / /api/stats /api/time /healthz\n") - task.OS.print("web_server PID: {}\n".format(task.getID())) - + listener = _open_listener(kernel, HOST, PORT, LISTEN_BACKLOG) try: + task.OS.print("smallOS web app running on http://{}:{}/\n".format(HOST, PORT)) + task.OS.print("routes: / /api/stats /api/time /healthz\n") + task.OS.print("web_server PID: {}\n".format(task.getID())) + while True: try: - client_sock, client_addr = listener.accept() + client_sock, client_addr = kernel.socket_accept(listener) except Exception as exc: - if kernel.socket_needs_read(exc): + retry_mode = kernel.socket_retry_mode(exc, "accept") + if retry_mode == "read": await task.wait_readable(listener) continue - if kernel.socket_needs_write(exc): + if retry_mode == "write": await task.wait_writable(listener) continue raise - kernel.socket_setblocking(client_sock, False) - task.spawn( - web_client_handler, - priority=max(1, task.priority - 1), - name="http_client", - args=(client_sock, client_addr, state), - ) + _dispatch_client(task, client_sock, client_addr, state) await task.yield_now() finally: - kernel.socket_close(listener) + _close_after_setup_failure(kernel, listener) def main(): diff --git a/tests/test_kernel.py b/tests/test_kernel.py index 3ccb98a..beeec31 100644 --- a/tests/test_kernel.py +++ b/tests/test_kernel.py @@ -4,7 +4,7 @@ sys.path.append("..") -from SmallPackage.Kernel import ESP32, ESP8266, MicroPythonKernel, PicoW, RaspberryPiPicoW, Unix, build_micropython_kernel +from SmallPackage.Kernel import Kernel, ESP32, ESP8266, MicroPythonKernel, PicoW, RaspberryPiPicoW, Unix, build_micropython_kernel class FakeNIC: @@ -172,7 +172,292 @@ class FakeSelectWithoutPoll: POLLOUT = 0x004 +class FailingWakeEndpoint: + def __init__(self, fail_setblocking=False, close_failures=0): + self.fail_setblocking = fail_setblocking + self.close_failures = close_failures + self.close_calls = 0 + self.closed = False + + def setblocking(self, _flag): + if self.fail_setblocking: + raise RuntimeError("setblocking failed") + + def close(self): + self.close_calls += 1 + if self.close_calls <= self.close_failures: + raise OSError("close failed") + self.closed = True + + +class FakeSocketPairModule: + def __init__(self, reader, writer): + self.reader = reader + self.writer = writer + + def socketpair(self): + return self.reader, self.writer + + +class TerminalWakeWriter: + def __init__(self, writer, *, interrupt=False, send_zero=False, close_failures=0): + self.writer = writer + self.interrupt = interrupt + self.send_zero = send_zero + self.close_failures = close_failures + self.send_calls = 0 + self.close_calls = 0 + + def setblocking(self, flag): + self.writer.setblocking(flag) + + def send(self, _data): + self.send_calls += 1 + if self.interrupt: + raise InterruptedError() + if self.send_zero: + return 0 + raise OSError("notification failed") + + def close(self): + self.close_calls += 1 + if self.close_calls <= self.close_failures: + raise OSError("temporary close failure") + self.writer.close() + + +class InvalidResultWakeWriter(TerminalWakeWriter): + def send(self, _data): + self.send_calls += 1 + return None + + +class InterruptingWakeReader(FailingWakeEndpoint): + def __init__(self): + super().__init__() + self.recv_calls = 0 + + def recv(self, _size): + self.recv_calls += 1 + raise InterruptedError() + + +class SendingWakeWriter(FailingWakeEndpoint): + def __init__(self): + super().__init__() + self.send_calls = 0 + + def send(self, data): + self.send_calls += 1 + return len(data) + + +class InvalidWakeReader(FailingWakeEndpoint): + pass + + +class InvalidWakeWriter(FailingWakeEndpoint): + pass + + class TestKernelProfiles(unittest.TestCase): + def test_base_and_micropython_wakeup_capabilities_are_explicit(self): + base = Kernel() + micropython = MicroPythonKernel() + micropython._socket = type( + "SocketPairPresent", + (), + {"socketpair": staticmethod(lambda: ())}, + )() + + for kernel in (base, micropython): + with self.subTest(kernel=type(kernel).__name__): + self.assertFalse(kernel.supports_wakeup_channel()) + with self.assertRaises(NotImplementedError): + kernel.create_wakeup_channel() + + def test_unix_wakeup_capability_requires_callable_socketpair(self): + kernel = Unix() + for socket_pair, expected in ( + (None, False), + (object(), False), + (lambda: (), True), + ): + with self.subTest(socket_pair=socket_pair): + kernel._socket = type("SocketModule", (), {"socketpair": socket_pair})() + self.assertEqual(expected, kernel.supports_wakeup_channel()) + if not expected: + with self.assertRaises(NotImplementedError): + kernel.create_wakeup_channel() + + def test_unix_wakeup_creation_cleans_every_acquired_endpoint(self): + reader = FailingWakeEndpoint(close_failures=1) + reader.recv = lambda _size: b"" + writer = SendingWakeWriter() + writer.fail_setblocking = True + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + + with self.assertRaisesRegex(RuntimeError, "setblocking failed"): + kernel.create_wakeup_channel() + + self.assertEqual(1, reader.close_calls) + self.assertEqual(1, writer.close_calls) + self.assertTrue(writer.closed) + + def test_unix_wakeup_creation_rejects_invalid_or_duplicate_endpoints(self): + cases = ( + (InvalidWakeReader(), SendingWakeWriter()), + (InterruptingWakeReader(), InvalidWakeWriter()), + ) + duplicate = SendingWakeWriter() + cases += ((duplicate, duplicate),) + + for reader, writer in cases: + with self.subTest(reader=type(reader).__name__, writer=type(writer).__name__): + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + with self.assertRaises((TypeError, ValueError)): + kernel.create_wakeup_channel() + self.assertEqual(1, reader.close_calls) + if writer is not reader: + self.assertEqual(1, writer.close_calls) + + def test_unix_wakeup_coalesces_and_reuses_notifications(self): + kernel = Unix() + channel = kernel.create_wakeup_channel() + try: + for _ in range(1000): + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + + channel.drain() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=0) + self.assertEqual([], readable) + + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + finally: + channel.close() + channel.close() + channel.notify() + channel.drain() + + def test_unix_wakeup_send_failure_becomes_readable_eof(self): + reader, raw_writer = socket.socketpair() + writer = TerminalWakeWriter(raw_writer) + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + channel = kernel.create_wakeup_channel() + wait_set = kernel.create_io_wait_set() + try: + wait_set.set_interest(channel.wait_object, True, False) + channel.notify() + channel.notify() + readable, _ = wait_set.wait(timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + self.assertEqual(1, writer.send_calls) + channel.drain() + finally: + wait_set.set_interest(channel.wait_object, False, False) + wait_set.close() + channel.close() + + def test_unix_wakeup_zero_send_becomes_readable_eof(self): + reader, raw_writer = socket.socketpair() + writer = TerminalWakeWriter(raw_writer, send_zero=True) + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + channel = kernel.create_wakeup_channel() + try: + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + self.assertEqual(1, writer.send_calls) + finally: + channel.close() + + def test_unix_wakeup_invalid_send_result_becomes_readable_eof(self): + reader, raw_writer = socket.socketpair() + writer = InvalidResultWakeWriter(raw_writer) + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + channel = kernel.create_wakeup_channel() + try: + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + self.assertEqual(1, writer.send_calls) + finally: + channel.close() + + def test_unix_wakeup_failed_terminal_close_remains_retryable(self): + reader, raw_writer = socket.socketpair() + writer = TerminalWakeWriter(raw_writer, close_failures=1) + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + channel = kernel.create_wakeup_channel() + try: + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=0) + self.assertEqual([], readable) + + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + self.assertEqual(2, writer.send_calls) + self.assertEqual(2, writer.close_calls) + finally: + channel.close() + + def test_unix_wakeup_bounds_interrupted_notify_and_drain(self): + reader, raw_writer = socket.socketpair() + writer = TerminalWakeWriter(raw_writer, interrupt=True) + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + channel = kernel.create_wakeup_channel() + try: + channel.notify() + readable, _ = kernel.io_wait([channel.wait_object], [], timeout_ms=100) + self.assertEqual([channel.wait_object], readable) + self.assertEqual(8, writer.send_calls) + finally: + channel.close() + + reader = InterruptingWakeReader() + writer = SendingWakeWriter() + kernel = Unix() + kernel._socket = FakeSocketPairModule(reader, writer) + channel = kernel.create_wakeup_channel() + try: + channel.notify() + channel.drain() + channel.notify() + self.assertEqual(8, reader.recv_calls) + self.assertEqual(1, writer.send_calls) + finally: + channel.close() + + def test_closed_unix_wakeup_detaches_from_persistent_wait_set(self): + kernel = Unix() + channel = kernel.create_wakeup_channel() + wait_set = kernel.create_io_wait_set() + wait_object = channel.wait_object + try: + wait_set.set_interest(wait_object, True, False) + channel.notify() + readable, _ = wait_set.wait(timeout_ms=100) + self.assertEqual([wait_object], readable) + + channel.close() + wait_set.set_interest(wait_object, False, False) + self.assertEqual(([], []), wait_set.wait(timeout_ms=0)) + finally: + channel.close() + wait_set.close() + def test_build_micropython_kernel_detects_esp32_profile(self): kernel = build_micropython_kernel(machine_name="ESP32 module with ESP32") diff --git a/tests/test_runtime.py b/tests/test_runtime.py index 288940e..c92ea7c 100644 --- a/tests/test_runtime.py +++ b/tests/test_runtime.py @@ -1,6 +1,7 @@ import os import sys import socket +import threading import unittest sys.path.append("..") @@ -425,6 +426,66 @@ async def writer(task, sock): self.assertEqual(b"x", reader_task.result) + def test_unix_wakeup_channel_wakes_scheduler_across_repeated_thread_cycles(self): + kernel = Unix() + channel = kernel.create_wakeup_channel() + wait_entered = threading.Semaphore(0) + original_create_wait_set = kernel.create_io_wait_set + + class ObservedWaitSet: + def __init__(self, inner): + self.inner = inner + self.readables = set() + + def set_interest(self, obj, readable, writable): + if readable: + self.readables.add(obj) + else: + self.readables.discard(obj) + return self.inner.set_interest(obj, readable, writable) + + def wait(self, timeout_ms=None): + if channel.wait_object in self.readables: + wait_entered.release() + return self.inner.wait(timeout_ms) + + def close(self): + return self.inner.close() + + kernel.create_io_wait_set = lambda: ObservedWaitSet(original_create_wait_set()) + runtime = SmallOS().setKernel(kernel) + detached = [] + notifier_errors = [] + + async def waiter(task): + for _ in range(3): + ready = await task.wait_readable(channel.wait_object) + detached.append(ready not in task.OS.ioReadWaiters) + channel.drain() + return "woke" + + def notify_scheduler(): + for _ in range(3): + if not wait_entered.acquire(timeout=1): + notifier_errors.append("scheduler did not enter persistent wait") + channel.notify() + + waiter_task = SmallTask(2, waiter, name="wakeup-waiter") + runtime.fork(waiter_task) + notifier = threading.Thread(target=notify_scheduler) + notifier.start() + try: + runtime.startOS() + finally: + notifier.join(timeout=1) + channel.close() + + self.assertFalse(notifier.is_alive()) + self.assertEqual([], notifier_errors) + self.assertEqual("woke", waiter_task.result) + self.assertEqual([True, True, True], detached) + self.assertNotIn(channel.wait_object, runtime.ioReadWaiters) + def test_killing_io_waiter_clears_wait_registration(self): """Cancelling an I/O waiter should remove it from the runtime waiter map.""" io_obj = object() diff --git a/tests/test_server_kernel.py b/tests/test_server_kernel.py new file mode 100644 index 0000000..4c5ef78 --- /dev/null +++ b/tests/test_server_kernel.py @@ -0,0 +1,430 @@ +import asyncio +import os +import socket +import sys +import unittest + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "demos")) + +from demos.web_app_demo import _dispatch_client, _open_listener, web_server_task +from SmallPackage.Kernel import Kernel, MicroPythonKernel, Unix + + +class FakeStream: + def __init__(self): + self.blocking_calls = [] + self.option_calls = [] + self.bind_calls = [] + self.listen_calls = [] + self.accept_result = (object(), ("192.0.2.2", 5000)) + self.local_address = ("0.0.0.0", 8080) + self.peer_address = ("192.0.2.2", 5000) + self.peer_error = None + self.close_calls = 0 + + def setblocking(self, flag): + self.blocking_calls.append(flag) + + def setsockopt(self, level, option, value): + self.option_calls.append((level, option, value)) + + def bind(self, address): + self.bind_calls.append(address) + + def listen(self, backlog): + self.listen_calls.append(backlog) + + def accept(self): + return self.accept_result + + def getsockname(self): + return self.local_address + + def getpeername(self): + if self.peer_error is not None: + raise self.peer_error + return self.peer_address + + def send(self, data): + return len(data) + + def recv(self, _size): + return b"" + + def close(self): + self.close_calls += 1 + + +class FakeSocketFactory: + setblocking = FakeStream.setblocking + setsockopt = FakeStream.setsockopt + bind = FakeStream.bind + listen = FakeStream.listen + accept = FakeStream.accept + getsockname = FakeStream.getsockname + getpeername = FakeStream.getpeername + send = FakeStream.send + recv = FakeStream.recv + close = FakeStream.close + + def __init__(self, module): + self.module = module + + def __call__(self, family, socktype, proto): + self.module.socket_calls.append((family, socktype, proto)) + return self.module.stream + + +class TwoArgumentSocketModule: + """MicroPython-style module exposing only a narrow resolver signature.""" + + SOCK_STREAM = 1 + + def __init__(self, stream): + self.stream = stream + self.resolve_calls = [] + self.socket_calls = [] + self.socket = FakeSocketFactory(self) + + def getaddrinfo(self, *args): + self.resolve_calls.append(args) + if len(args) != 2: + raise TypeError("this port accepts only host and port") + host, port = args + return [(2, self.SOCK_STREAM, 6, "", (host, port))] + + +class MissingServerSocketModule: + SOCK_STREAM = 1 + + +class IncompleteStream: + def __init__(self): + self.close_calls = 0 + + def close(self): + self.close_calls += 1 + + +class RecordingServerKernel: + def __init__(self, fail_at=None, close_fails=False, supported=True): + self.fail_at = fail_at + self.close_fails = close_fails + self.supported = supported + self.calls = [] + self.record = object() + self.listener = object() + + def _call(self, name): + self.calls.append(name) + if self.fail_at == name: + raise RuntimeError("{} failed".format(name)) + + def supports_tcp_server(self): + self.calls.append("supports_tcp_server") + return self.supported + + def supports_reuse_address(self): + self.calls.append("supports_reuse_address") + return True + + def resolve_passive_address(self, host, port): + self._call("resolve_passive_address") + self.resolved = (host, port) + return self.record + + def socket_open(self, record): + self._call("socket_open") + self.open_record = record + return self.listener + + def socket_set_reuse_address(self, listener, enabled): + self._call("socket_set_reuse_address") + self.reuse_args = (listener, enabled) + + def socket_bind(self, listener, record): + self._call("socket_bind") + self.bind_args = (listener, record) + + def socket_listen(self, listener, backlog): + self._call("socket_listen") + self.listen_args = (listener, backlog) + + def socket_setblocking(self, listener, flag): + self._call("socket_setblocking") + self.blocking_args = (listener, flag) + + def socket_accept(self, listener): + self._call("socket_accept") + return object(), ("192.0.2.1", 5000) + + def socket_retry_mode(self, exc, operation): + return None + + def socket_close(self, listener): + self.calls.append("socket_close") + self.closed_listener = listener + if self.close_fails: + raise RuntimeError("close failed") + + +class DispatchTask: + priority = 2 + + def __init__(self, kernel, spawn_fails=False): + self.OS = type("OS", (), {"kernel": kernel})() + self.spawn_fails = spawn_fails + self.spawn_calls = [] + + def spawn(self, routine, **kwargs): + self.spawn_calls.append((routine, kwargs)) + if self.spawn_fails: + raise RuntimeError("spawn failed") + + +class ServerTask(DispatchTask): + def __init__(self, kernel, output_fails=False): + super().__init__(kernel) + self.output_fails = output_fails + self.OS.print = self._print + + def _print(self, _message): + if self.output_fails: + raise RuntimeError("output failed") + + def getID(self): + return 1 + + +class TestServerKernelContract(unittest.TestCase): + def test_base_kernel_reports_unsupported_passive_operations(self): + kernel = Kernel() + + self.assertFalse(kernel.supports_tcp_server()) + self.assertFalse(kernel.supports_reuse_address()) + unsupported_calls = ( + lambda: kernel.resolve_passive_address("", 0), + lambda: kernel.socket_open(object()), + lambda: kernel.socket_setblocking(object(), False), + lambda: kernel.socket_set_reuse_address(object(), True), + lambda: kernel.socket_bind(object(), object()), + lambda: kernel.socket_listen(object(), 1), + lambda: kernel.socket_accept(object()), + lambda: kernel.socket_local_address(object()), + lambda: kernel.socket_peer_address(object()), + lambda: kernel.socket_close(object()), + ) + for call in unsupported_calls: + with self.subTest(call=call), self.assertRaises(NotImplementedError): + call() + + def test_unix_loopback_echo_uses_only_kernel_operations(self): + kernel = Unix() + record = kernel.resolve_passive_address("127.0.0.1", 0) + listener = kernel.socket_open(record) + client = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + accepted = None + try: + self.assertTrue(kernel.supports_tcp_server()) + self.assertNotIsInstance(record, tuple) + if kernel.supports_reuse_address(): + kernel.socket_set_reuse_address(listener, True) + try: + kernel.socket_bind(listener, record) + except PermissionError as exc: + self.skipTest("loopback bind is not permitted: {}".format(exc)) + kernel.socket_listen(listener, 2) + local_address = kernel.socket_local_address(listener) + self.assertGreater(local_address[1], 0) + + client.connect(local_address) + accepted, accepted_address = kernel.socket_accept(listener) + self.assertEqual(client.getsockname(), accepted_address) + self.assertEqual(accepted_address, kernel.socket_peer_address(accepted)) + self.assertEqual(local_address, kernel.socket_local_address(accepted)) + + client.sendall(b"request") + self.assertEqual(b"request", kernel.socket_recv(accepted, 7)) + self.assertEqual(8, kernel.socket_send(accepted, b"response")) + self.assertEqual(b"response", client.recv(8)) + finally: + if accepted is not None: + kernel.socket_close(accepted) + kernel.socket_close(listener) + client.close() + + def test_unix_nonblocking_accept_exhaustion_is_readable_retry(self): + kernel = Unix() + record = kernel.resolve_passive_address("127.0.0.1", 0) + listener = kernel.socket_open(record) + try: + try: + kernel.socket_bind(listener, record) + except PermissionError as exc: + self.skipTest("loopback bind is not permitted: {}".format(exc)) + kernel.socket_listen(listener, 1) + kernel.socket_setblocking(listener, False) + with self.assertRaises(BlockingIOError) as raised: + kernel.socket_accept(listener) + self.assertEqual( + "read", kernel.socket_retry_mode(raised.exception, "accept") + ) + finally: + kernel.socket_close(listener) + + def test_unix_rejects_passive_record_from_another_kernel(self): + first = Unix() + second = Unix() + record = first.resolve_passive_address("127.0.0.1", 0) + + with self.assertRaises(ValueError): + second.socket_open(record) + with self.assertRaises(ValueError): + second.socket_bind(object(), record) + + def test_peer_address_suppresses_only_not_connected(self): + unix = Unix() + listener = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + try: + self.assertIsNone(unix.socket_peer_address(listener)) + finally: + listener.close() + with self.assertRaises(OSError): + unix.socket_peer_address(listener) + + class FakeErrno: + ENOTCONN = 57 + + stream = FakeStream() + socket_module = TwoArgumentSocketModule(stream) + kernel = MicroPythonKernel( + modules={"socket": socket_module, "errno": FakeErrno()} + ) + stream.peer_error = OSError(57, "not connected") + self.assertIsNone(kernel.socket_peer_address(stream)) + stream.peer_error = OSError(9, "bad descriptor") + with self.assertRaises(OSError): + kernel.socket_peer_address(stream) + + def test_micropython_complete_port_uses_opaque_record_and_constants(self): + stream = FakeStream() + socket_module = TwoArgumentSocketModule(stream) + socket_module.SOL_SOCKET = 7 + socket_module.SO_REUSEADDR = 9 + kernel = MicroPythonKernel(modules={"socket": socket_module}) + + self.assertTrue(kernel.supports_tcp_server()) + self.assertTrue(kernel.supports_reuse_address()) + record = kernel.resolve_passive_address("", 8080) + self.assertNotIsInstance(record, tuple) + self.assertEqual(2, len(socket_module.resolve_calls)) + self.assertEqual(("0.0.0.0", 8080), socket_module.resolve_calls[-1]) + self.assertIs(stream, kernel.socket_open(record)) + + kernel.socket_set_reuse_address(stream, True) + kernel.socket_bind(stream, record) + kernel.socket_listen(stream, 7) + self.assertEqual([(7, 9, 1)], stream.option_calls) + self.assertEqual([("0.0.0.0", 8080)], stream.bind_calls) + self.assertEqual([7], stream.listen_calls) + self.assertEqual(stream.accept_result, kernel.socket_accept(stream)) + self.assertEqual(stream.local_address, kernel.socket_local_address(stream)) + self.assertEqual(stream.peer_address, kernel.socket_peer_address(stream)) + + def test_micropython_does_not_guess_missing_reuse_constants(self): + kernel = MicroPythonKernel( + modules={"socket": TwoArgumentSocketModule(FakeStream())} + ) + + self.assertTrue(kernel.supports_tcp_server()) + self.assertFalse(kernel.supports_reuse_address()) + with self.assertRaises(NotImplementedError): + kernel.socket_set_reuse_address(FakeStream(), True) + + def test_micropython_incomplete_port_reports_unsupported(self): + kernel = MicroPythonKernel(modules={"socket": MissingServerSocketModule()}) + + self.assertFalse(kernel.supports_tcp_server()) + with self.assertRaises(NotImplementedError): + kernel.resolve_passive_address("", 80) + + def test_micropython_incomplete_opened_stream_closes_once(self): + stream = IncompleteStream() + socket_module = TwoArgumentSocketModule(stream) + kernel = MicroPythonKernel(modules={"socket": socket_module}) + + with self.assertRaises(NotImplementedError): + kernel.socket_open(kernel.resolve_passive_address("", 80)) + self.assertEqual(1, stream.close_calls) + + def test_listener_uses_same_record_for_open_and_bind(self): + kernel = RecordingServerKernel() + + listener = _open_listener(kernel, "127.0.0.1", 0, 5) + + self.assertIs(kernel.listener, listener) + self.assertIs(kernel.record, kernel.open_record) + self.assertIs(kernel.record, kernel.bind_args[1]) + self.assertEqual((kernel.listener, False), kernel.blocking_args) + + def test_listener_capability_failure_precedes_resolution_and_binding(self): + kernel = RecordingServerKernel(supported=False) + + with self.assertRaises(NotImplementedError): + _open_listener(kernel, "", 80, 1) + + self.assertEqual(["supports_tcp_server"], kernel.calls) + + def test_listener_post_open_failures_close_once_and_preserve_primary_error(self): + setup_steps = ( + "socket_set_reuse_address", + "socket_bind", + "socket_listen", + "socket_setblocking", + ) + for step in setup_steps: + with self.subTest(step=step): + kernel = RecordingServerKernel(fail_at=step, close_fails=True) + with self.assertRaisesRegex(RuntimeError, "{} failed".format(step)): + _open_listener(kernel, "", 80, 1) + self.assertEqual(1, kernel.calls.count("socket_close")) + self.assertIs(kernel.listener, kernel.closed_listener) + + def test_accepted_stream_failures_close_once_and_preserve_primary_error(self): + for fail_at, spawn_fails, message in ( + ("socket_setblocking", False, "socket_setblocking failed"), + (None, True, "spawn failed"), + ): + with self.subTest(message=message): + kernel = RecordingServerKernel( + fail_at=fail_at, + close_fails=True, + ) + task = DispatchTask(kernel, spawn_fails=spawn_fails) + stream = object() + with self.assertRaisesRegex(RuntimeError, message): + _dispatch_client(task, stream, ("192.0.2.1", 5000), {}) + self.assertEqual(1, kernel.calls.count("socket_close")) + self.assertIs(stream, kernel.closed_listener) + + def test_server_startup_failure_closes_listener_once(self): + kernel = RecordingServerKernel(close_fails=True) + task = ServerTask(kernel, output_fails=True) + + with self.assertRaisesRegex(RuntimeError, "output failed"): + asyncio.run(web_server_task(task, {})) + + self.assertEqual(1, kernel.calls.count("socket_close")) + + def test_server_loop_error_is_not_masked_by_close_failure(self): + kernel = RecordingServerKernel(fail_at="socket_accept", close_fails=True) + task = ServerTask(kernel) + + with self.assertRaisesRegex(RuntimeError, "socket_accept failed"): + asyncio.run(web_server_task(task, {})) + + self.assertEqual(1, kernel.calls.count("socket_close")) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_socket_retry.py b/tests/test_socket_retry.py new file mode 100644 index 0000000..e65eed2 --- /dev/null +++ b/tests/test_socket_retry.py @@ -0,0 +1,214 @@ +import asyncio +import errno +import os +import sys +import unittest + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "demos")) + +from demos.web_app_demo import _read_request_head, _send_all +from SmallPackage.Kernel import Kernel, MicroPythonKernel, Unix +from SmallPackage.clients.SmallStream import SmallStream + + +class FakeSSLWantReadError(OSError): + pass + + +class FakeSSLWantWriteError(OSError): + pass + + +class FakeSSLModule: + SSLWantReadError = FakeSSLWantReadError + SSLWantWriteError = FakeSSLWantWriteError + + +class StreamSocket: + pass + + +class ConnectWouldBlockSocket: + def connect(self, _sockaddr): + raise BlockingIOError() + + +class StreamKernel(Kernel): + def __init__(self): + super().__init__() + self.send_attempts = [] + self.recv_attempts = 0 + self.handshake_attempts = 0 + + def resolve_address(self, host, port): + return (0, 0, 0, "", (host, port)) + + def socket_open(self, address_info): + return StreamSocket() + + def socket_setblocking(self, sock, flag): + return None + + def socket_connect(self, sock, sockaddr): + return True + + def socket_wrap_tls_client(self, sock, **kwargs): + return sock + + def socket_do_handshake(self, sock): + self.handshake_attempts += 1 + if self.handshake_attempts == 1: + raise BlockingIOError(errno.EAGAIN, "try again") + + def socket_send(self, sock, data): + self.send_attempts.append(data) + if len(self.send_attempts) == 1: + raise BlockingIOError(errno.EAGAIN, "try again") + return len(data) + + def socket_recv(self, sock, buffer_size): + self.recv_attempts += 1 + if self.recv_attempts == 1: + raise BlockingIOError(errno.EAGAIN, "try again") + return b"response" + + +class DemoStreamKernel(StreamKernel): + def socket_recv(self, sock, buffer_size): + self.recv_attempts += 1 + if self.recv_attempts == 1: + raise BlockingIOError(errno.EAGAIN, "try again") + return b"GET / HTTP/1.1\r\n\r\n" + + +class RuntimeRef: + def __init__(self, kernel): + self.kernel = kernel + + +class StreamTask: + def __init__(self, kernel): + self.OS = RuntimeRef(kernel) + self.read_waits = [] + self.write_waits = [] + + async def wait_readable(self, sock): + self.read_waits.append(sock) + + async def wait_writable(self, sock): + self.write_waits.append(sock) + + +class TestSocketRetryClassification(unittest.TestCase): + def test_base_retry_direction_depends_on_operation(self): + kernel = Kernel() + error = BlockingIOError(errno.EAGAIN, "try again") + + self.assertEqual("read", kernel.socket_retry_mode(error, "accept")) + self.assertEqual("read", kernel.socket_retry_mode(error, "recv")) + self.assertEqual("read", kernel.socket_retry_mode(error, "handshake")) + self.assertEqual("write", kernel.socket_retry_mode(error, "send")) + self.assertIsNone(kernel.socket_retry_mode(OSError(errno.EINVAL, "bad"), "recv")) + + def test_unix_retry_mode_preserves_tls_direction(self): + kernel = Unix() + pending = BlockingIOError(errno.EAGAIN, "try again") + + self.assertEqual("read", kernel.socket_retry_mode(pending, "recv")) + self.assertEqual("write", kernel.socket_retry_mode(pending, "send")) + self.assertEqual( + "write", kernel.socket_retry_mode(OSError(errno.EAGAIN, "again"), "send") + ) + self.assertEqual( + "read", + kernel.socket_retry_mode(kernel._ssl.SSLWantReadError(), "send"), + ) + self.assertEqual( + "write", + kernel.socket_retry_mode(kernel._ssl.SSLWantWriteError(), "recv"), + ) + self.assertIsNone( + kernel.socket_retry_mode( + BlockingIOError(errno.EINPROGRESS, "pending connect"), "send" + ) + ) + + def test_micropython_matches_unix_and_preserves_tls_direction(self): + kernel = MicroPythonKernel(modules={"ssl": FakeSSLModule()}) + pending = OSError(errno.EAGAIN, "try again") + + self.assertEqual("read", kernel.socket_retry_mode(pending, "accept")) + self.assertEqual("read", kernel.socket_retry_mode(pending, "recv")) + self.assertEqual("write", kernel.socket_retry_mode(pending, "send")) + self.assertEqual( + "read", kernel.socket_retry_mode(FakeSSLWantReadError(), "send") + ) + self.assertEqual( + "write", kernel.socket_retry_mode(FakeSSLWantWriteError(), "recv") + ) + self.assertIsNone( + kernel.socket_retry_mode(OSError(115, "pending connect"), "recv") + ) + + def test_micropython_connect_preserves_errno_less_would_block(self): + kernel = MicroPythonKernel() + + self.assertFalse(kernel.socket_connect(ConnectWouldBlockSocket(), ("host", 80))) + + def test_unknown_operation_is_rejected_before_classification(self): + kernels = (Kernel(), Unix(), MicroPythonKernel()) + + for kernel in kernels: + with self.subTest(kernel=type(kernel).__name__): + with self.assertRaisesRegex(ValueError, "Unknown socket operation"): + kernel.socket_retry_mode(BlockingIOError(), "connect") + + def test_small_stream_forwards_memoryview_and_waits_by_operation(self): + kernel = StreamKernel() + task = StreamTask(kernel) + stream = SmallStream(task, "example.test", 80) + stream.sock = StreamSocket() + stream._connected = True + payload = memoryview(b"request") + + async def exercise_stream(): + await stream.send_all(payload) + return await stream.recv_some() + + result = asyncio.run(exercise_stream()) + + self.assertEqual(b"response", result) + self.assertIs(payload, kernel.send_attempts[0]) + self.assertEqual([stream.sock], task.write_waits) + self.assertEqual([stream.sock], task.read_waits) + + def test_small_stream_handshake_waits_for_readability(self): + kernel = StreamKernel() + task = StreamTask(kernel) + stream = SmallStream(task, "example.test", 443, use_tls=True) + + result = asyncio.run(stream.connect()) + + self.assertIs(stream, result) + self.assertEqual(2, kernel.handshake_attempts) + self.assertEqual([stream.sock], task.read_waits) + self.assertEqual([], task.write_waits) + + def test_web_demo_routes_send_and_recv_backpressure_by_operation(self): + kernel = DemoStreamKernel() + task = StreamTask(kernel) + sock = StreamSocket() + + async def exercise_helpers(): + await _send_all(task, sock, memoryview(b"response")) + return await _read_request_head(task, sock) + + result = asyncio.run(exercise_helpers()) + + self.assertEqual(b"GET / HTTP/1.1\r\n\r\n", result) + self.assertEqual([sock], task.write_waits) + self.assertEqual([sock], task.read_waits) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/typing/kernel_socket_contract.py b/tests/typing/kernel_socket_contract.py new file mode 100644 index 0000000..caa7e9d --- /dev/null +++ b/tests/typing/kernel_socket_contract.py @@ -0,0 +1,34 @@ +# pyright: strict, reportUnnecessaryTypeIgnoreComment=true + +"""Static contract for operation-aware retries and bytes-like socket sends.""" + +from typing import Any + +from SmallPackage.Kernel import Kernel +from SmallPackage._types import SocketBuffer, SocketKernelLike, SocketOperation + + +class SocketKernel(Kernel): + received: SocketBuffer | None = None + + def socket_send(self, sock: Any, data: SocketBuffer) -> int: + self.received = data + return len(data) + + +def accepts_kernel(kernel: SocketKernelLike) -> None: + """Require the concrete kernel to satisfy the public structural contract.""" + + +def verify_socket_contract() -> None: + kernel = SocketKernel() + accepts_kernel(kernel) + payload = memoryview(b"response") + sent: int = kernel.socket_send(object(), payload) + operation: SocketOperation = "send" + retry_mode = kernel.socket_retry_mode(BlockingIOError(), operation) + assert sent == len(payload) + assert kernel.received is payload + assert retry_mode == "write" + + kernel.socket_retry_mode(BlockingIOError(), "connect") # pyright: ignore[reportArgumentType] diff --git a/tests/typing/passive_tcp_kernel_contract.py b/tests/typing/passive_tcp_kernel_contract.py new file mode 100644 index 0000000..ad31a3c --- /dev/null +++ b/tests/typing/passive_tcp_kernel_contract.py @@ -0,0 +1,10 @@ +"""Static consumer contract for passive TCP kernels.""" + +from SmallPackage.Kernel import Unix +from SmallPackage._types import PassiveTCPKernelLike + + +kernel: PassiveTCPKernelLike = Unix() +address = kernel.resolve_passive_address("127.0.0.1", 0) +listener = kernel.socket_open(address) +kernel.socket_bind(listener, address)