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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
73 changes: 69 additions & 4 deletions SmallPackage/Kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,21 @@
if TYPE_CHECKING:
from collections.abc import Iterable, Mapping, Sequence
from typing import Any, cast
from ._types import SocketBuffer, SocketOperation, SocketRetryMode


_UNSET = object()
_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)
)
)


def _import_first(*module_names: str) -> Any | None:
Expand Down Expand Up @@ -556,7 +568,8 @@ 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:
Expand Down Expand Up @@ -593,6 +606,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.
Expand Down Expand Up @@ -766,7 +790,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):
Expand Down Expand Up @@ -811,6 +835,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):
'''
Expand Down Expand Up @@ -848,6 +891,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')
Expand Down Expand Up @@ -957,7 +1001,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

Expand All @@ -966,7 +1010,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):
Expand Down Expand Up @@ -1050,6 +1094,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.
Expand Down
14 changes: 13 additions & 1 deletion SmallPackage/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -42,6 +45,15 @@ def io_wait(
) -> tuple[Sequence[Any], Sequence[Any]]: ...


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."""

Expand Down
17 changes: 10 additions & 7 deletions SmallPackage/clients/SmallStream.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -135,18 +136,19 @@ 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)
if sent == 0:
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
Expand All @@ -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
Expand Down
17 changes: 10 additions & 7 deletions demos/web_app_demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -369,10 +371,11 @@ async def web_server_task(task, state):
try:
client_sock, client_addr = listener.accept()
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
Expand Down
Loading
Loading