diff --git a/README.md b/README.md index f3b5caf..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 diff --git a/SmallPackage/Kernel.py b/SmallPackage/Kernel.py index 24fd224..2a8a2b6 100644 --- a/SmallPackage/Kernel.py +++ b/SmallPackage/Kernel.py @@ -24,10 +24,12 @@ 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') @@ -396,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. @@ -539,6 +652,14 @@ 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 @@ -691,6 +812,7 @@ def __init__(self): import socket import ssl import sys + import threading import time self._errno = errno @@ -708,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 @@ -804,6 +927,69 @@ 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 @@ -1079,6 +1265,16 @@ 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 diff --git a/SmallPackage/_types.py b/SmallPackage/_types.py index 5326b2b..0c38612 100644 --- a/SmallPackage/_types.py +++ b/SmallPackage/_types.py @@ -45,6 +45,24 @@ 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.""" 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()