diff --git a/README.md b/README.md index 08f7d73..48bf207 100644 --- a/README.md +++ b/README.md @@ -103,10 +103,87 @@ async def hello(task): config = SmallOSConfig.from_json_file("smallos.config.json") runtime = SmallOS(config=config).setKernel(Unix()) +runtime.setErrorHandler( + lambda event: print( + "[smallOS] task failure in {} (PID {}): {}".format( + event["task_name"] or "unnamed task", + event["task_id"], + event["exception_repr"], + ) + ) +) runtime.fork([SmallTask(2, hello, name="hello")]) runtime.startOS() ``` +## Runtime Error Handling + +`smallOS` now supports a runtime-level error observer through +`runtime.setErrorHandler(handler, include_cancelled=False)`. + +Use it when you want: +- readable debug output for uncaught task failures +- lightweight cleanup or bookkeeping at the runtime boundary +- a single place to surface task errors without crashing the scheduler + +The handler is synchronous and receives a failure-event dictionary after the +task has been finalized. Current event fields include: +- `task_id` +- `task_name` +- `parent_id` +- `exception` +- `exception_type` +- `exception_repr` +- `is_cancelled` +- `blocked_reason` +- `waiting_signal` +- `io_wait_mode` +- `join_target_id` +- `join_pending_ids` +- `traceback_text` + +By default, `TaskCancelledError` does not trigger the handler. Pass +`include_cancelled=True` if you want cancellation events too. + +Example: + +```python +from SmallPackage.Kernel import Unix +from SmallPackage.SmallOS import SmallOS + + +def log_runtime_error(event): + print( + "[smallOS] task failure in {} (PID {}): {}".format( + event["task_name"] or "unnamed task", + event["task_id"], + event["exception_repr"], + ) + ) + if event["traceback_text"]: + print(event["traceback_text"], end="") + + +runtime = SmallOS().setKernel(Unix()) +runtime.setErrorHandler(log_runtime_error) +``` + +### Closed or Invalid File Descriptors + +Closed or invalid file descriptors used in `wait_readable(...)` or +`wait_writable(...)` no longer crash the whole scheduler through the platform +poll/select layer. + +Instead: +- the kernel validates the watched object before polling +- the waiting task receives a normal exception such as `ValueError` +- the runtime finalizes that task cleanly +- your runtime error handler can log or clean up the failure gracefully + +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. + ## Configuration The runtime now uses a first-class config object backed by @@ -201,6 +278,12 @@ The new demos live in [demos](demos): - [demos/mqtt_demo.py](demos/mqtt_demo.py): MQTT example built on the native cooperative client +All of the shared demo entry points now install a default runtime error handler +through [demos/common.py](demos/common.py). That means network failures, +invalid I/O wait objects, and other uncaught task exceptions are reported as +readable task-failure diagnostics instead of looking like abrupt scheduler +crashes or silent exits. + The original root demo remains available in [demo.py](demo.py) as a compatibility wrapper around [demos/runtime_demo.py](demos/runtime_demo.py). diff --git a/SmallPackage/Kernel.py b/SmallPackage/Kernel.py index c3c0d16..f1e3176 100644 --- a/SmallPackage/Kernel.py +++ b/SmallPackage/Kernel.py @@ -90,6 +90,270 @@ def _portable_shell_split(line): return tokens +def _io_wait_lookup_key(obj): + """Return the stable descriptor/key used by poll-style backends.""" + if hasattr(obj, 'fileno'): + try: + return obj.fileno() + except Exception: + pass + return obj + + +class _SelectorIOWaitSet: + """Persistent CPython selector with delta-updated registrations.""" + + def __init__(self, selector, read_mask, write_mask): + self._selector = selector + self._read_mask = read_mask + self._write_mask = write_mask + self._entries = {} + self._objects_by_key = {} + self._closed = False + + def _event_mask(self, readable, writable): + mask = 0 + if readable: + mask |= self._read_mask + if writable: + mask |= self._write_mask + return mask + + def _remove(self, obj): + entry = self._entries.pop(obj, None) + if entry is None: + return + _, key = entry + if self._objects_by_key.get(key) is obj: + del self._objects_by_key[key] + try: + self._selector.unregister(key) + except (KeyError, OSError, ValueError): + # The descriptor may already have been closed and removed by the OS. + pass + + def _discard_stale_key_owner(self, obj, key): + existing_obj = self._objects_by_key.get(key) + if existing_obj is None or existing_obj is obj: + return + + if _io_wait_lookup_key(existing_obj) == key: + raise ValueError( + 'I/O descriptor {} is already registered by another object.'.format(key) + ) + self._remove(existing_obj) + + def set_interest(self, obj, readable, writable): + """Apply one object's combined read/write interest when it changes.""" + if self._closed: + raise RuntimeError('I/O wait set is closed.') + + mask = self._event_mask(readable, writable) + entry = self._entries.get(obj) + if entry is None: + if mask == 0: + return + lookup_key = _io_wait_lookup_key(obj) + self._discard_stale_key_owner(obj, lookup_key) + selector_key = self._selector.register(obj, mask, obj) + key = selector_key.fd + self._entries[obj] = (mask, key) + self._objects_by_key[key] = obj + return + + old_mask, key = entry + if mask == 0: + self._remove(obj) + return + if mask == old_mask: + return + + current_key = _io_wait_lookup_key(obj) + if current_key != key: + self._remove(obj) + self.set_interest(obj, readable, writable) + return + + self._selector.modify(key, mask, obj) + self._entries[obj] = (mask, key) + + def wait(self, timeout_ms=None): + """Wait for readiness without rebuilding the registered descriptor set.""" + if self._closed: + raise RuntimeError('I/O wait set is closed.') + timeout = None if timeout_ms is None else max(0, timeout_ms) / 1000 + events = self._selector.select(timeout) + ready_read = [] + ready_write = [] + for selector_key, mask in events: + obj = selector_key.data + entry = self._entries.get(obj) + if entry is None or entry[1] != selector_key.fd: + raise RuntimeError('I/O selector returned a stale readiness event.') + if mask & self._read_mask: + ready_read.append(obj) + if mask & self._write_mask: + ready_write.append(obj) + return ready_read, ready_write + + def close(self): + """Release selector resources without closing registered user objects.""" + if self._closed: + return + for obj in list(self._entries): + self._remove(obj) + self._selector.close() + self._closed = True + + +class _PollIOWaitSet: + """Portable persistent poll set used by MicroPython kernels.""" + + def __init__(self, select_mod, poller): + self._poller = poller + self._read_mask = getattr(select_mod, 'POLLIN', 0x001) + self._write_mask = getattr(select_mod, 'POLLOUT', 0x004) + self._invalid_mask = getattr(select_mod, 'POLLNVAL', 0x020) + self._error_mask = ( + getattr(select_mod, 'POLLERR', 0x008) + | getattr(select_mod, 'POLLHUP', 0x010) + | getattr(select_mod, 'POLLRDHUP', 0) + | self._invalid_mask + ) + self._entries = {} + self._objects_by_key = {} + self._closed = False + + def _event_mask(self, readable, writable): + mask = 0 + if readable: + mask |= self._read_mask + if writable: + mask |= self._write_mask + return mask + + def _remove(self, obj): + entry = self._entries.pop(obj, None) + if entry is None: + return + _, key = entry + if self._objects_by_key.get(key) is obj: + del self._objects_by_key[key] + + try: + self._poller.unregister(obj) + return + except Exception: + pass + try: + self._poller.unregister(key) + except Exception: + # Some MicroPython ports discard closed streams automatically. + pass + + def _discard_stale_key_owner(self, obj, key): + existing_obj = self._objects_by_key.get(key) + if existing_obj is None or existing_obj is obj: + return + if _io_wait_lookup_key(existing_obj) == key: + raise ValueError( + 'I/O descriptor {} is already registered by another object.'.format(key) + ) + self._remove(existing_obj) + + def set_interest(self, obj, readable, writable): + """Apply one poll registration delta.""" + if self._closed: + raise RuntimeError('I/O wait set is closed.') + + mask = self._event_mask(readable, writable) + entry = self._entries.get(obj) + if entry is None: + if mask == 0: + return + key = _io_wait_lookup_key(obj) + self._discard_stale_key_owner(obj, key) + self._poller.register(obj, mask) + self._entries[obj] = (mask, key) + self._objects_by_key[key] = obj + return + + old_mask, key = entry + if mask == 0: + self._remove(obj) + return + if mask == old_mask: + return + + current_key = _io_wait_lookup_key(obj) + if current_key != key: + self._remove(obj) + self.set_interest(obj, readable, writable) + return + + modifier = getattr(self._poller, 'modify', None) + if modifier is not None: + modifier(obj, mask) + else: + # MicroPython permits repeated register() calls to update the mask. + self._poller.register(obj, mask) + self._entries[obj] = (mask, key) + + def _ready_object(self, event_obj): + try: + if event_obj in self._entries: + return event_obj + except (KeyError, TypeError): + pass + try: + return self._objects_by_key[event_obj] + except (KeyError, TypeError): + raise RuntimeError('poll returned an unknown I/O object.') + + def wait(self, timeout_ms=None): + """Consume poll/ipoll events immediately and return original objects.""" + if self._closed: + raise RuntimeError('I/O wait set is closed.') + timeout = -1 if timeout_ms is None else max(0, int(timeout_ms)) + ipoll = getattr(self._poller, 'ipoll', None) + events = ipoll(timeout) if callable(ipoll) else self._poller.poll(timeout) + + ready_read = [] + ready_write = [] + for event in events: + event_obj = event[0] + mask = event[1] + obj = self._ready_object(event_obj) + entry = self._entries.get(obj) + if entry is None: + raise RuntimeError('poll returned a stale readiness event.') + requested_mask = entry[0] + ready_mask = mask + if mask & self._error_mask: + # HUP/ERR apply to the requested directions even though callers do + # not include those unsolicited bits in the registration mask. + ready_mask |= requested_mask + if ready_mask & self._read_mask: + ready_read.append(obj) + if ready_mask & self._write_mask: + ready_write.append(obj) + if mask & self._invalid_mask: + self._remove(obj) + return ready_read, ready_write + + def close(self): + """Unregister streams and clear constrained-runtime references.""" + if self._closed: + return + for obj in list(self._entries): + self._remove(obj) + closer = getattr(self._poller, 'close', None) + if closer is not None: + closer() + self._objects_by_key.clear() + self._closed = True + + def detect_micropython_machine_name(sys_mod=None, os_mod=None): """ Best-effort lookup of the active board/firmware machine name. @@ -218,6 +482,37 @@ def sleep_ms(self, delay_ms): def io_wait(self, readables, writables, timeout_ms=None): return [], [] + def create_io_wait_set(self): + """Return an optional persistent readiness set for one scheduler run.""" + return None + + def validate_io_wait_object(self, obj): + """ + Return whether ``obj`` still looks safe to hand to poll/select. + + Closed sockets on CPython typically report ``fileno() == -1`` after close. + Letting those reach ``select.poll().register(...)`` crashes the runtime + with ``ValueError`` before the waiting task can be resumed or cleaned up. + """ + if obj is None: + return False, ValueError('I/O wait object cannot be None.') + if isinstance(obj, int): + if obj < 0: + return False, ValueError( + 'I/O wait object has invalid file descriptor ({}).'.format(obj) + ) + return True, None + if hasattr(obj, 'fileno'): + try: + fd = obj.fileno() + except Exception as exc: + return False, exc + if isinstance(fd, int) and fd < 0: + return False, ValueError( + 'I/O wait object has invalid file descriptor ({}).'.format(fd) + ) + return True, None + def resolve_address(self, host, port): return None @@ -278,12 +573,7 @@ def _poll_lookup_key(self, obj): original socket object, so the kernel needs a stable way to translate poll results back into the scheduler's waiter keys. """ - if hasattr(obj, 'fileno'): - try: - return obj.fileno() - except Exception: - pass - return obj + return _io_wait_lookup_key(obj) class Unix(Kernel): @@ -298,16 +588,27 @@ class Unix(Kernel): def __init__(self): super().__init__() import errno + import os import shlex import select + import selectors import socket import ssl import sys import time self._errno = errno + self._os = os self._shlex = shlex self._select = select + self._selector_factory = selectors.DefaultSelector + if sys.platform == 'darwin' and hasattr(selectors, 'PollSelector'): + # Kqueue has a comparatively high fixed cost for zero-time checks on + # macOS. Persistent poll avoids the small-wait-set scheduler regression + # while still removing registration rebuilds. + self._selector_factory = selectors.PollSelector + self._selector_read_mask = selectors.EVENT_READ + self._selector_write_mask = selectors.EVENT_WRITE self._socket = socket self._ssl = ssl self._sys = sys @@ -381,6 +682,29 @@ def io_wait(self, readables, writables, timeout_ms=None): ready_read, ready_write, _ = self._select.select(readables, writables, [], timeout) return ready_read, ready_write + def create_io_wait_set(self): + """Create the runtime-owned selector used for persistent Unix waits.""" + return _SelectorIOWaitSet( + self._selector_factory(), + self._selector_read_mask, + self._selector_write_mask, + ) + + def validate_io_wait_object(self, obj): + """Reject Unix descriptors that were closed behind a persistent wait set.""" + is_valid, exc = super().validate_io_wait_object(obj) + if not is_valid: + return is_valid, exc + + fd = _io_wait_lookup_key(obj) + if not isinstance(fd, int): + return True, None + try: + self._os.fstat(fd) + except (OSError, OverflowError, ValueError) as exc: + return False, exc + return True, None + def resolve_address(self, host, port): return self._socket.getaddrinfo(host, port, type=self._socket.SOCK_STREAM)[0] @@ -566,6 +890,12 @@ def io_wait(self, readables, writables, timeout_ms=None): self.sleep_ms(timeout_ms) return [], [] + def create_io_wait_set(self): + """Reuse one poller when the active MicroPython port supplies it.""" + if self._poll_factory is None: + return None + return _PollIOWaitSet(self._select, self._poll_factory()) + def resolve_address(self, host, port): return self._socket.getaddrinfo(host, port)[0] diff --git a/SmallPackage/OSlist.py b/SmallPackage/OSlist.py index dc60ab5..b2e02c8 100644 --- a/SmallPackage/OSlist.py +++ b/SmallPackage/OSlist.py @@ -118,6 +118,18 @@ def pop(self): return task return None + def has_ready(self): + """Return whether a valid runnable task is queued without removing it.""" + for priority in range(1, self.num_priorities): + queue = self.ready[priority] + while queue: + task = queue[0] + if self.search(task.getID()) != -1 and task.getExeStatus(): + return True + queue.popleft() + task._queued = False + return False + def add_sleeping(self, task, wake_time): """Push a sleeping task onto the wake-time heap.""" self._sleep_seq += 1 diff --git a/SmallPackage/SmallOS.py b/SmallPackage/SmallOS.py index 043d1eb..45fd7a7 100644 --- a/SmallPackage/SmallOS.py +++ b/SmallPackage/SmallOS.py @@ -15,7 +15,7 @@ from .SmallIO import SmallIO from .SmallConfig import SmallOSConfig from .OSlist import OSList -from .SmallErrors import MaxProcessError, UnsupportedAwaitableError +from .SmallErrors import MaxProcessError, TaskCancelledError, UnsupportedAwaitableError _MISSING = object() @@ -58,11 +58,14 @@ def __init__(self, size=None, config=None, **kwargs): self.wakeUpdate = [] self.ioReadWaiters = {} self.ioWriteWaiters = {} + self._io_wait_set = None self.shells = [] self.tasks = OSList(self.config.priority_levels, self.config.task_capacity) self.kernel = None self.eternalWatchers = self.config.eternal_watchers self.cursor = None + self.errorHandler = None + self.errorHandlerIncludeCancelled = False SmallIO.__init__(self, self.config.io_buffer_length) if kwargs: @@ -89,27 +92,68 @@ def start(self): task, advances it once, and then either finalizes it or handles the wait condition it requested. """ - while len(self.tasks) != 0: - self._wake_sleeping_tasks() - self._wake_io_tasks(timeout_ms=0) - self.cursor = self.tasks.pop() - - if self.cursor is None: - if not self._idle_until_next_task(): - break - continue + self._open_io_wait_set() + try: + while len(self.tasks) != 0: + self._wake_sleeping_tasks() + if self.tasks.has_ready(): + # Include newly ready I/O tasks in the priority decision. + self._wake_io_tasks(timeout_ms=0) + else: + # Go directly to the blocking wait instead of first issuing + # a redundant zero-time poll. + if not self._idle_until_next_task(): + break + self._wake_sleeping_tasks() + + self.cursor = self.tasks.pop() + if self.cursor is None: + continue - yielded = self.cursor.execute() - if self.cursor.done: - # Finished tasks are finalized immediately so PID lookup and join - # bookkeeping always see a consistent terminal state. - self._finalize_task(self.cursor) - else: - self._handle_yield(self.cursor, yielded) + yielded = self.cursor.execute() + if self.cursor.done: + # Finished tasks are finalized immediately so PID lookup and join + # bookkeeping always see a consistent terminal state. + self._finalize_task(self.cursor) + else: + self._handle_yield(self.cursor, yielded) - if not self.eternalWatchers and len(self.tasks) != 0 and self.tasks.isOnlyWatchers(): - return - return + if not self.eternalWatchers and len(self.tasks) != 0 and self.tasks.isOnlyWatchers(): + return + return + finally: + self._close_io_wait_set() + + def _open_io_wait_set(self): + """Create and seed the kernel's optional persistent readiness set.""" + self._close_io_wait_set() + if self.kernel is None: + return + factory = getattr(self.kernel, "create_io_wait_set", None) + if factory is None: + return + wait_set = factory() + if wait_set is None: + return + + self._io_wait_set = wait_set + try: + self._fail_invalid_io_waiters() + for io_obj in self.ioReadWaiters: + self._refresh_io_interest(io_obj) + for io_obj in self.ioWriteWaiters: + if io_obj not in self.ioReadWaiters: + self._refresh_io_interest(io_obj) + except BaseException: + self._close_io_wait_set() + raise + + def _close_io_wait_set(self): + """Close only the backend wait set, never the registered user objects.""" + wait_set = self._io_wait_set + self._io_wait_set = None + if wait_set is not None: + wait_set.close() def next(self): """Return the next runnable task without advancing the main loop.""" @@ -146,6 +190,20 @@ def setEternalWatchers(self, isEternalWatcherPresent): self.eternalWatchers = isEternalWatcherPresent return self + def setErrorHandler(self, handler, include_cancelled=False): + """ + Install a best-effort runtime error observer for failed tasks. + + ``handler`` is a synchronous callable that receives one failure-event + dictionary after the task has been finalized. Returning ``None`` removes + the current handler. + """ + if handler is not None and not callable(handler): + raise TypeError("error handler must be callable or None.") + self.errorHandler = handler + self.errorHandlerIncludeCancelled = bool(include_cancelled) + return self + def _wake_sleeping_tasks(self): """Promote every expired sleeping task back onto the ready queues.""" if not self.kernel: @@ -184,7 +242,7 @@ def resume_task(self, task, value=_MISSING, exc=None, front=False): """ Requeue a blocked task with the value or exception that completed it. - Centralizing resume logic here ensures stale join registrations are + Centralizing resume logic here ensures scheduler-owned wait state is cleared before a task gets a chance to block on something else. """ if task is None or task == -1 or task.done: @@ -192,7 +250,7 @@ def resume_task(self, task, value=_MISSING, exc=None, front=False): if self.tasks.search(task.getID()) == -1: return -1 - self._clear_wait_registration(task) + self._clear_wait_state(task) task.resume(value=value, exc=exc) self.tasks.enqueue(task, front=front) return 0 @@ -243,9 +301,7 @@ def _handle_yield(self, task, yielded): delay_ms = max(0, int(seconds * 1000)) wake_time = self.kernel.scheduler_now_ms() + delay_ms if self.kernel else delay_ms - task.block("sleep") - task._wake_at = wake_time - self.tasks.add_sleeping(task, wake_time) + self._enter_sleep_wait(task, wake_time) return if operation == "wait_signal": @@ -253,22 +309,15 @@ def _handle_yield(self, task, yielded): if task.checkSignal(signal): self.resume_task(task, value=signal, front=True) else: - task.block("signal") - task._waiting_signal = signal + self._enter_signal_wait(task, signal) return if operation == "wait_readable": - task.block("wait_readable") - task._io_wait_obj = payload["io_obj"] - task._io_wait_mode = "read" - self._register_io_wait(task, payload["io_obj"], "read") + self._enter_io_wait(task, payload["io_obj"], "read") return if operation == "wait_writable": - task.block("wait_writable") - task._io_wait_obj = payload["io_obj"] - task._io_wait_mode = "write" - self._register_io_wait(task, payload["io_obj"], "write") + self._enter_io_wait(task, payload["io_obj"], "write") return if operation == "join": @@ -281,9 +330,7 @@ def _handle_yield(self, task, yielded): if target.done: self._resume_from_completed(task, target) else: - task.block("join") - task._join_target = target - target.add_join_waiter(task) + self._enter_join_wait(task, target) return if operation == "join_all": @@ -309,12 +356,7 @@ def _handle_yield(self, task, yielded): # Preserve the original child ordering for the eventual results # while also tracking a fast set of outstanding child PIDs. - task.block("join_all") - task._join_targets = targets - task._join_pending = pending - for child in targets: - if not child.done: - child.add_join_waiter(task) + self._enter_join_all_wait(task, targets, pending) return task.fail(UnsupportedAwaitableError("Unknown instruction {!r}".format(operation))) @@ -354,13 +396,104 @@ def _first_exception(self, tasks): return task.exception return None + def _clear_wait_state(self, task): + """ + Remove a task from scheduler wait bookkeeping and reset its wait metadata. + + The scheduler owns the blocked-state lifecycle, so runtime transitions + clear both registration-based wait structures and the task's stored wait + markers from one place. + """ + self._clear_wait_registration(task) + self._clear_wait_metadata(task) + + def _clear_wait_metadata(self, task): + """Reset the scheduler-owned wait metadata stored on ``task``.""" + task._blocked_reason = None + task._wake_at = None + task._waiting_signal = None + task._join_target = None + task._join_targets = None + task._join_pending = set() + task._io_wait_obj = None + task._io_wait_mode = None + + def _begin_wait(self, task, reason): + """Prepare a runnable task to transition into one blocked wait state.""" + self._clear_wait_state(task) + task.block(reason) + + def _enter_sleep_wait(self, task, wake_time): + """Put ``task`` to sleep until ``wake_time``.""" + self._begin_wait(task, "sleep") + task._wake_at = wake_time + self.tasks.add_sleeping(task, wake_time) + + def _enter_signal_wait(self, task, signal): + """Block ``task`` until ``signal`` is delivered.""" + self._begin_wait(task, "signal") + task._waiting_signal = signal + + def _enter_io_wait(self, task, io_obj, mode): + """Register ``task`` for readable or writable I/O readiness.""" + reason = "wait_readable" if mode == "read" else "wait_writable" + self._begin_wait(task, reason) + task._io_wait_obj = io_obj + task._io_wait_mode = mode + validator = getattr(self.kernel, "validate_io_wait_object", None) + if validator is not None: + is_valid, exc = validator(io_obj) + if not is_valid: + self.resume_task(task, exc=self._clone_wait_error(exc), front=True) + return + try: + self._register_io_wait(task, io_obj, mode) + except Exception as exc: + # Registration failures belong at the await expression; they should + # not tear down the entire scheduler loop. + self.resume_task(task, exc=exc, front=True) + + def _enter_join_wait(self, task, target): + """Block ``task`` until ``target`` finishes.""" + self._begin_wait(task, "join") + task._join_target = target + target.add_join_waiter(task) + + def _enter_join_all_wait(self, task, targets, pending): + """Block ``task`` until every child in ``pending`` has completed.""" + self._begin_wait(task, "join_all") + task._join_targets = list(targets) + task._join_pending = set(pending) + for child in targets: + if not child.done: + child.add_join_waiter(task) + def _register_io_wait(self, task, io_obj, mode): """Register a task as waiting on an I/O object's readiness event.""" waiters = self.ioReadWaiters if mode == "read" else self.ioWriteWaiters + added_interest = io_obj not in waiters if io_obj not in waiters: waiters[io_obj] = [] if task not in waiters[io_obj]: waiters[io_obj].append(task) + if added_interest: + try: + self._refresh_io_interest(io_obj) + except Exception: + waiters[io_obj].remove(task) + if not waiters[io_obj]: + del waiters[io_obj] + raise + + def _refresh_io_interest(self, io_obj): + """Apply one object's combined logical interest to a persistent wait set.""" + if self._io_wait_set is None: + return + self._io_wait_set.set_interest( + io_obj, + io_obj in self.ioReadWaiters, + io_obj in self.ioWriteWaiters, + ) def _wake_io_tasks(self, timeout_ms=0): """ @@ -374,19 +507,81 @@ def _wake_io_tasks(self, timeout_ms=0): return if not self.ioReadWaiters and not self.ioWriteWaiters: return + self._fail_invalid_io_waiters() + if not self.ioReadWaiters and not self.ioWriteWaiters: + return - readable, writable = self.kernel.io_wait( - list(self.ioReadWaiters.keys()), - list(self.ioWriteWaiters.keys()), - timeout_ms, - ) - self._resume_io_waiters(readable, self.ioReadWaiters) - self._resume_io_waiters(writable, self.ioWriteWaiters) + if self._io_wait_set is not None: + readable, writable = self._io_wait_set.wait(timeout_ms) + else: + readable, writable = self.kernel.io_wait( + list(self.ioReadWaiters.keys()), + list(self.ioWriteWaiters.keys()), + timeout_ms, + ) + self._resume_ready_io(readable, writable) + + def _fail_invalid_io_waiters(self): + """ + Resume waiters whose I/O objects are already closed or otherwise invalid. + + Poll/select backends raise immediately when handed a stale descriptor, + which would otherwise take down the whole runtime before the affected + task can observe the problem. + """ + validator = getattr(self.kernel, "validate_io_wait_object", None) + if validator is None: + return + io_objects = list(self.ioReadWaiters.keys()) + for io_obj in self.ioWriteWaiters: + if io_obj not in self.ioReadWaiters: + io_objects.append(io_obj) + + for io_obj in io_objects: + is_valid, exc = validator(io_obj) + if is_valid: + continue - def _resume_io_waiters(self, ready_objects, waiters_map): - """Resume every task waiting on the now-ready I/O objects.""" - for io_obj in ready_objects: - waiters = waiters_map.pop(io_obj, []) + waiters = self.ioReadWaiters.pop(io_obj, []) + for waiter in self.ioWriteWaiters.pop(io_obj, []): + if waiter not in waiters: + waiters.append(waiter) + self._refresh_io_interest(io_obj) + for waiter in waiters: + if waiter.done or self.tasks.search(waiter.getID()) == -1: + continue + self.resume_task(waiter, exc=self._clone_wait_error(exc), front=True) + + def _clone_wait_error(self, exc): + """Return a fresh exception instance for resuming a blocked waiter.""" + if isinstance(exc, BaseException): + args = getattr(exc, "args", ()) + try: + return exc.__class__(*args) + except Exception: + return RuntimeError(str(exc)) + return RuntimeError("I/O wait object is no longer valid.") + + def _resume_ready_io(self, readable, writable): + """Detach one readiness snapshot before applying registration deltas.""" + batches = [] + changed_objects = [] + for ready_objects, waiters_map in ( + (readable, self.ioReadWaiters), + (writable, self.ioWriteWaiters), + ): + for io_obj in ready_objects: + waiters = waiters_map.pop(io_obj, []) + batches.append((io_obj, waiters)) + if io_obj not in changed_objects: + changed_objects.append(io_obj) + + # An object ready for both directions should move directly from its + # combined mask to its final mask instead of modify-then-unregister. + for io_obj in changed_objects: + self._refresh_io_interest(io_obj) + + for io_obj, waiters in batches: for waiter in waiters: if waiter.done or self.tasks.search(waiter.getID()) == -1: continue @@ -394,30 +589,106 @@ def _resume_io_waiters(self, ready_objects, waiters_map): def _clear_wait_registration(self, task): """ - Remove a task from any join bookkeeping it currently participates in. + Remove a task from any registration-based wait bookkeeping. This prevents leaked waiter references when a task is resumed, cancelled, or moved from one wait condition to another. """ if task._join_target is not None: task._join_target.discard_join_waiter(task) - task._join_target = None if task._join_targets: for child in task._join_targets: child.discard_join_waiter(task) - task._join_targets = None - task._join_pending = set() if task._io_wait_obj is not None and task._io_wait_mode is not None: + io_obj = task._io_wait_obj waiters_map = self.ioReadWaiters if task._io_wait_mode == "read" else self.ioWriteWaiters - waiters = waiters_map.get(task._io_wait_obj, []) + waiters = waiters_map.get(io_obj, []) while task in waiters: waiters.remove(task) - if not waiters and task._io_wait_obj in waiters_map: - del waiters_map[task._io_wait_obj] - task._io_wait_obj = None - task._io_wait_mode = None + if not waiters and io_obj in waiters_map: + del waiters_map[io_obj] + self._refresh_io_interest(io_obj) + + def _should_dispatch_failure(self, task): + """Return whether ``task`` should produce a runtime failure event.""" + exc = task.exception + if exc is None: + return False + if isinstance(exc, TaskCancelledError) and not self.errorHandlerIncludeCancelled: + return False + return True + + def _snapshot_task_id(self, task): + """Return a task or PID reference as a PID integer when possible.""" + if task is None or task == -1: + return None + if hasattr(task, "getID"): + return task.getID() + if isinstance(task, int): + return task + return None + + def _format_exception_traceback(self, exc): + """Return a best-effort formatted traceback string for ``exc``.""" + try: + import traceback + except ImportError: + return None + + try: + return "".join( + traceback.format_exception(type(exc), exc, getattr(exc, "__traceback__", None)) + ) + except Exception: + try: + return "".join(traceback.format_exception_only(type(exc), exc)) + except Exception: + return None + + def _build_failure_event(self, task): + """Snapshot the task failure context before finalization clears wait state.""" + exc = task.exception + return { + "task_id": task.getID(), + "task_name": task.name, + "parent_id": self._snapshot_task_id(task.parent), + "exception": exc, + "exception_type": type(exc).__name__ if exc is not None else None, + "exception_repr": repr(exc), + "is_cancelled": isinstance(exc, TaskCancelledError), + "blocked_reason": task._blocked_reason, + "waiting_signal": task._waiting_signal, + "io_wait_mode": task._io_wait_mode, + "join_target_id": self._snapshot_task_id(task._join_target), + "join_pending_ids": sorted(task._join_pending) if task._join_pending else [], + "traceback_text": self._format_exception_traceback(exc), + } + + def _write_runtime_diagnostic(self, message): + """Write a best-effort runtime diagnostic without crashing the scheduler.""" + if not message: + return + if not message.endswith("\n"): + message += "\n" + try: + if self.kernel and hasattr(self.kernel, "write"): + self.kernel.write(message) + except Exception: + return + + def _dispatch_error_handler(self, event): + """Invoke the installed runtime error handler without surfacing its failures.""" + if self.errorHandler is None: + return + try: + self.errorHandler(event) + except Exception as exc: + diagnostic = self._format_exception_traceback(exc) or repr(exc) + self._write_runtime_diagnostic( + "smallOS error handler failed: {}".format(diagnostic.rstrip("\n")) + ) def _detach_from_parent(self, task): """Remove a finished child PID from its parent's child list.""" @@ -445,7 +716,6 @@ def _notify_waiters(self, task): continue if waiter._blocked_reason == "join" and waiter._join_target is task: - waiter._join_target = None self._resume_from_completed(waiter, task) continue @@ -453,22 +723,25 @@ def _notify_waiters(self, task): if task.exception is not None: # ``join_all`` behaves like structured concurrency here: # one child failure wakes the parent immediately. - self._clear_wait_registration(waiter) self.resume_task(waiter, exc=task.exception, front=True) continue waiter._join_pending.discard(task.getID()) if not waiter._join_pending: results = [child.result for child in waiter._join_targets] - self._clear_wait_registration(waiter) self.resume_task(waiter, value=results, front=True) def _finalize_task(self, task): """Run the full shutdown sequence for a finished or cancelled task.""" - self._clear_wait_registration(task) + failure_event = None + if self._should_dispatch_failure(task): + failure_event = self._build_failure_event(task) + self._clear_wait_state(task) self._notify_waiters(task) self._detach_from_parent(task) self.tasks.delete(task.getID()) + if failure_event is not None: + self._dispatch_error_handler(failure_event) def cancel_task(self, task, recursive=False): """Cancel a task by object or PID and optionally cancel its descendants.""" diff --git a/SmallPackage/SmallTask.py b/SmallPackage/SmallTask.py index 62e5247..2a88215 100644 --- a/SmallPackage/SmallTask.py +++ b/SmallPackage/SmallTask.py @@ -197,11 +197,6 @@ def complete(self, result): """Mark the task as successfully finished and store its result.""" self._done = True self._result = result - self._blocked_reason = None - self._wake_at = None - self._waiting_signal = None - self._io_wait_obj = None - self._io_wait_mode = None self.isReady = 0 self.isWaiting = 0 self.isSleep = 0 @@ -212,11 +207,6 @@ def fail(self, exc): """Mark the task as failed and store its terminal exception.""" self._done = True self._exception = exc - self._blocked_reason = None - self._wake_at = None - self._waiting_signal = None - self._io_wait_obj = None - self._io_wait_mode = None self.isReady = 0 self.isWaiting = 0 self.isSleep = 0 @@ -245,13 +235,9 @@ def resume(self, value=_MISSING, exc=None): Prepare the task to run again after a wait condition completes. ``value`` is sent into the coroutine on the next step. ``exc`` is thrown - into it instead. The scheduler chooses which of those channels to use. + into it instead. The scheduler chooses which of those channels to use + and is responsible for clearing any wait metadata beforehand. """ - self._blocked_reason = None - self._wake_at = None - self._waiting_signal = None - self._io_wait_obj = None - self._io_wait_mode = None self.isReady = 1 self.isWaiting = 0 self.isSleep = 0 diff --git a/benchmarks/io_wait_benchmark.py b/benchmarks/io_wait_benchmark.py new file mode 100644 index 0000000..b17ae83 --- /dev/null +++ b/benchmarks/io_wait_benchmark.py @@ -0,0 +1,196 @@ +#!/usr/bin/env python3 +"""Manual snapshot-versus-persistent Unix readiness benchmark. + +Performance assertions deliberately stay out of CI. Run this script on each +supported Unix/Python combination and retain the Python version, selector name, +descriptor count, readiness pattern, and medians with review notes. +""" + +import argparse +import os +import socket +import statistics +import sys +import time + + +REPOSITORY_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +if REPOSITORY_ROOT not in sys.path: + sys.path.insert(0, REPOSITORY_ROOT) + +from SmallPackage.Kernel import Unix +from SmallPackage.SmallOS import SmallOS +from SmallPackage.SmallTask import SmallTask + + +class SnapshotUnix(Unix): + """Unix kernel forced through the compatibility snapshot path.""" + + def create_io_wait_set(self): + return None + + +def _elapsed_microseconds(callable_obj, iterations): + started = time.perf_counter() + for _ in range(iterations): + callable_obj() + return (time.perf_counter() - started) * 1000000 / iterations + + +def _median_microseconds(callable_obj, iterations, repeats): + samples = [ + _elapsed_microseconds(callable_obj, iterations) + for _ in range(repeats) + ] + return statistics.median(samples) + + +def _prime_readiness(socket_pairs, readiness): + if readiness == "idle": + return + selected = socket_pairs[:1] if readiness == "one-ready" else socket_pairs + for _reader, writer in selected: + writer.send(b"x") + + +def benchmark_case(descriptor_count, readiness, iterations, repeats): + socket_pairs = [socket.socketpair() for _ in range(descriptor_count)] + readers = [pair[0] for pair in socket_pairs] + kernel = Unix() + wait_set = kernel.create_io_wait_set() + try: + for reader in readers: + reader.setblocking(False) + wait_set.set_interest(reader, True, False) + _prime_readiness(socket_pairs, readiness) + + snapshot = _median_microseconds( + lambda: kernel.io_wait(readers, [], 0), + iterations, + repeats, + ) + persistent = _median_microseconds( + lambda: wait_set.wait(0), + iterations, + repeats, + ) + return snapshot, persistent + finally: + wait_set.close() + for reader, writer in socket_pairs: + reader.close() + writer.close() + + +async def _idle_io_watcher(task, reader): + await task.wait_readable(reader) + + +async def _runnable_task(task, iterations): + for _ in range(iterations): + await task.yield_now() + + +def _scheduler_run(kernel_class, iterations): + reader, writer = socket.socketpair() + runtime = SmallOS().setKernel(kernel_class()) + watcher = SmallTask( + 1, + _idle_io_watcher, + name="idle-io-watcher", + args=(reader,), + isWatcher=True, + ) + runnable = SmallTask( + 2, + _runnable_task, + name="runnable-task", + args=(iterations,), + ) + runtime.fork([watcher, runnable]) + started = time.perf_counter() + try: + runtime.startOS() + return (time.perf_counter() - started) * 1000000 / iterations + finally: + reader.close() + writer.close() + + +def benchmark_scheduler(iterations, repeats): + snapshot_samples = [ + _scheduler_run(SnapshotUnix, iterations) + for _ in range(repeats) + ] + persistent_samples = [ + _scheduler_run(Unix, iterations) + for _ in range(repeats) + ] + return statistics.median(snapshot_samples), statistics.median(persistent_samples) + + +def _parse_counts(value): + return [int(item.strip()) for item in value.split(",") if item.strip()] + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--counts", default="1,4,32,256") + parser.add_argument("--iterations", type=int, default=2000) + parser.add_argument("--repeats", type=int, default=5) + parser.add_argument("--scheduler-iterations", type=int, default=10000) + args = parser.parse_args() + + probe = Unix().create_io_wait_set() + selector_name = type(probe._selector).__name__ + probe.close() + print("Python {} | {}".format(sys.version.split()[0], selector_name)) + print("fds readiness snapshot_us persistent_us speedup") + + for descriptor_count in _parse_counts(args.counts): + for readiness in ("idle", "one-ready", "all-ready"): + try: + snapshot, persistent = benchmark_case( + descriptor_count, + readiness, + args.iterations, + args.repeats, + ) + except OSError as exc: + print( + "{:4d} {:11s} skipped: {}".format( + descriptor_count, + readiness, + exc, + ) + ) + break + speedup = snapshot / persistent if persistent else float("inf") + print( + "{:4d} {:11s} {:11.2f} {:14.2f} {:7.2f}x".format( + descriptor_count, + readiness, + snapshot, + persistent, + speedup, + ) + ) + + snapshot, persistent = benchmark_scheduler( + args.scheduler_iterations, + args.repeats, + ) + speedup = snapshot / persistent if persistent else float("inf") + print() + print("Runnable task + one idle I/O watcher (microseconds/scheduler step)") + print( + "snapshot {:.2f} | persistent {:.2f} | {:.2f}x".format( + snapshot, + persistent, + speedup, + ) + ) + + +if __name__ == "__main__": + main() diff --git a/demos/common.py b/demos/common.py index ccaafe7..0b5063d 100644 --- a/demos/common.py +++ b/demos/common.py @@ -2,8 +2,8 @@ Shared demo helpers for desktop and board-specific smallOS examples. These helpers keep the individual demo files short while still showing the -recommended public API: load a config file, choose a kernel, spawn tasks, and -start the runtime. +recommended public API: load a config file, choose a kernel, install an error +handler, spawn tasks, and start the runtime. """ import os @@ -34,7 +34,50 @@ def load_demo_config(**overrides): def build_runtime(kernel, **config_overrides): """Create a ``SmallOS`` instance wired to the chosen kernel.""" - return SmallOS(config=load_demo_config(**config_overrides)).setKernel(kernel) + runtime = SmallOS(config=load_demo_config(**config_overrides)).setKernel(kernel) + return install_demo_error_handler(runtime) + + +def _format_failure_event(event): + """Return a readable multi-line summary for demo task failures.""" + details = [] + if event["parent_id"] is not None: + details.append("parent={}".format(event["parent_id"])) + if event["blocked_reason"] is not None: + details.append("blocked={}".format(event["blocked_reason"])) + if event["waiting_signal"] is not None: + details.append("signal={}".format(event["waiting_signal"])) + if event["io_wait_mode"] is not None: + details.append("io={}".format(event["io_wait_mode"])) + if event["join_target_id"] is not None: + details.append("join_target={}".format(event["join_target_id"])) + if event["join_pending_ids"]: + details.append("join_pending={}".format(event["join_pending_ids"])) + + header = "[smallOS demo] task failure" + if event["task_name"]: + header += " in {}".format(event["task_name"]) + if event["task_id"] is not None: + header += " (PID {})".format(event["task_id"]) + header += ": {}".format(event["exception_repr"]) + + if details: + header += " [{}]".format(", ".join(details)) + + trace = event.get("traceback_text") + if trace: + return "{}\n{}".format(header, trace if trace.endswith("\n") else trace + "\n") + return header + "\n" + + +def install_demo_error_handler(runtime, include_cancelled=False): + """Attach the shared demo error logger to ``runtime``.""" + + def _handler(event): + runtime.kernel.write(_format_failure_event(event)) + + runtime.setErrorHandler(_handler, include_cancelled=include_cancelled) + return runtime async def worker(task): diff --git a/tests/test_OSlist.py b/tests/test_OSlist.py index 5fb031b..60ffa2f 100644 --- a/tests/test_OSlist.py +++ b/tests/test_OSlist.py @@ -41,6 +41,20 @@ def test_search_and_delete(self): self.assertEqual(0, tasks.delete(pid)) self.assertEqual(-1, tasks.search(pid)) + def test_has_ready_preserves_live_task_and_discards_stale_entries(self): + tasks = OSList(10) + stale = SmallTask(1, None, name="stale") + live = SmallTask(2, None, name="live") + for task in (stale, live): + tasks.insert(task) + tasks.enqueue(task) + tasks.delete(stale.getID()) + + self.assertTrue(tasks.has_ready()) + self.assertFalse(stale._queued) + self.assertIs(live, tasks.pop()) + self.assertFalse(tasks.has_ready()) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_kernel.py b/tests/test_kernel.py index 86b84f7..3ccb98a 100644 --- a/tests/test_kernel.py +++ b/tests/test_kernel.py @@ -1,4 +1,5 @@ import sys +import socket import unittest sys.path.append("..") @@ -71,6 +72,106 @@ def poll(self, timeout): return list(self.events) +class FakeSelectorKey: + def __init__(self, fileobj, fd, events, data): + self.fileobj = fileobj + self.fd = fd + self.events = events + self.data = data + + +class CountingSelector: + def __init__(self): + self.keys = {} + self.events = [] + self.registrations = [] + self.modifications = [] + self.unregistrations = [] + self.timeouts = [] + self.closed = False + + def register(self, obj, events, data=None): + fd = obj if isinstance(obj, int) else obj.fileno() + key = FakeSelectorKey(obj, fd, events, data) + self.keys[fd] = key + self.registrations.append((obj, events, data)) + return key + + def modify(self, obj, events, data=None): + fd = obj if isinstance(obj, int) else obj.fileno() + previous = self.keys[fd] + key = FakeSelectorKey(previous.fileobj, fd, events, data) + self.keys[fd] = key + self.modifications.append((obj, events, data)) + return key + + def unregister(self, obj): + fd = obj if isinstance(obj, int) else obj.fileno() + self.unregistrations.append(fd) + return self.keys.pop(fd) + + def select(self, timeout=None): + self.timeouts.append(timeout) + return [(self.keys[fd], mask) for fd, mask in self.events] + + def close(self): + self.closed = True + + +class PersistentFakePoller: + def __init__(self, use_ipoll=True): + self.events = [] + self.registrations = [] + self.modifications = [] + self.unregistrations = [] + self.ipoll_timeouts = [] + self.poll_timeouts = [] + self.closed = False + if not use_ipoll: + self.ipoll = None + + def register(self, obj, mask): + self.registrations.append((obj, mask)) + + def modify(self, obj, mask): + self.modifications.append((obj, mask)) + + def unregister(self, obj): + self.unregistrations.append(obj) + + def ipoll(self, timeout): + self.ipoll_timeouts.append(timeout) + return iter(self.events) + + def poll(self, timeout): + self.poll_timeouts.append(timeout) + return list(self.events) + + def close(self): + self.closed = True + + +class FakeSelectModule: + POLLIN = 0x001 + POLLOUT = 0x004 + POLLERR = 0x008 + POLLHUP = 0x010 + POLLNVAL = 0x020 + + def __init__(self, poller=None): + self.poller = poller + self.poll_calls = 0 + + def poll(self): + self.poll_calls += 1 + return self.poller + + +class FakeSelectWithoutPoll: + POLLIN = 0x001 + POLLOUT = 0x004 + + class TestKernelProfiles(unittest.TestCase): def test_build_micropython_kernel_detects_esp32_profile(self): kernel = build_micropython_kernel(machine_name="ESP32 module with ESP32") @@ -148,6 +249,190 @@ def test_unix_io_wait_maps_poll_file_descriptors_back_to_objects(self): self.assertEqual([readable_obj], readable) self.assertEqual([writable_obj], writable) + def test_unix_wait_set_applies_only_registration_deltas(self): + readable_obj = FakePollObject(11) + selector = CountingSelector() + kernel = Unix() + kernel._selector_factory = lambda: selector + wait_set = kernel.create_io_wait_set() + + wait_set.set_interest(readable_obj, True, False) + wait_set.set_interest(readable_obj, True, False) + wait_set.set_interest(readable_obj, True, True) + selector.events = [(11, kernel._selector_read_mask | kernel._selector_write_mask)] + readable, writable = wait_set.wait(timeout_ms=5) + wait_set.set_interest(readable_obj, False, False) + wait_set.close() + + self.assertEqual(1, len(selector.registrations)) + self.assertEqual(1, len(selector.modifications)) + self.assertEqual([11], selector.unregistrations) + self.assertEqual([0.005], selector.timeouts) + self.assertEqual([readable_obj], readable) + self.assertEqual([readable_obj], writable) + self.assertTrue(selector.closed) + + def test_unix_wait_set_does_not_reregister_stable_objects(self): + selector = CountingSelector() + kernel = Unix() + kernel._selector_factory = lambda: selector + wait_set = kernel.create_io_wait_set() + io_objects = [FakePollObject(fd) for fd in range(100, 132)] + try: + for io_obj in io_objects: + wait_set.set_interest(io_obj, True, False) + for _ in range(100): + wait_set.wait(timeout_ms=0) + + self.assertEqual(32, len(selector.registrations)) + self.assertEqual([], selector.modifications) + self.assertEqual([], selector.unregistrations) + self.assertEqual(100, len(selector.timeouts)) + finally: + wait_set.close() + + def test_unix_wait_set_replaces_a_stale_descriptor_owner(self): + selector = CountingSelector() + kernel = Unix() + kernel._selector_factory = lambda: selector + wait_set = kernel.create_io_wait_set() + original = FakePollObject(140) + replacement = FakePollObject(140) + try: + wait_set.set_interest(original, True, False) + original.fd = -1 + wait_set.set_interest(replacement, True, False) + + self.assertEqual(2, len(selector.registrations)) + self.assertEqual([140], selector.unregistrations) + selector.events = [(140, kernel._selector_read_mask)] + readable, _writable = wait_set.wait(timeout_ms=0) + self.assertEqual([replacement], readable) + finally: + wait_set.close() + + def test_unix_wait_set_uses_real_socket_readiness_without_closing_socket(self): + left, right = socket.socketpair() + wait_set = Unix().create_io_wait_set() + try: + wait_set.set_interest(left, True, False) + right.send(b"x") + + readable, writable = wait_set.wait(timeout_ms=100) + self.assertEqual([left], readable) + self.assertEqual([], writable) + finally: + wait_set.close() + self.assertGreaterEqual(left.fileno(), 0) + left.close() + right.close() + + def test_unix_wait_set_reports_peer_hangup_as_readable(self): + left, right = socket.socketpair() + wait_set = Unix().create_io_wait_set() + try: + wait_set.set_interest(left, True, False) + right.close() + + readable, _writable = wait_set.wait(timeout_ms=100) + self.assertEqual([left], readable) + finally: + wait_set.close() + left.close() + + def test_micropython_wait_set_reuses_poller_and_ipoll(self): + poller = PersistentFakePoller() + select_mod = FakeSelectModule(poller) + kernel = MicroPythonKernel(modules={"select": select_mod}) + wait_set = kernel.create_io_wait_set() + io_obj = FakePollObject(21) + + wait_set.set_interest(io_obj, True, False) + wait_set.set_interest(io_obj, True, False) + wait_set.set_interest(io_obj, True, True) + poller.events = [(21, select_mod.POLLIN, "port-specific-extra")] + readable, writable = wait_set.wait(timeout_ms=7) + wait_set.set_interest(io_obj, False, False) + wait_set.close() + + self.assertEqual(1, select_mod.poll_calls) + self.assertEqual([(io_obj, select_mod.POLLIN)], poller.registrations) + self.assertEqual( + [(io_obj, select_mod.POLLIN | select_mod.POLLOUT)], + poller.modifications, + ) + self.assertEqual([io_obj], poller.unregistrations) + self.assertEqual([7], poller.ipoll_timeouts) + self.assertEqual([], poller.poll_timeouts) + self.assertEqual([io_obj], readable) + self.assertEqual([], writable) + self.assertTrue(poller.closed) + + def test_micropython_wait_set_falls_back_to_poll_and_maps_errors(self): + poller = PersistentFakePoller(use_ipoll=False) + select_mod = FakeSelectModule(poller) + kernel = MicroPythonKernel(modules={"select": select_mod}) + wait_set = kernel.create_io_wait_set() + readable_obj = FakePollObject(31) + writable_obj = FakePollObject(32) + + wait_set.set_interest(readable_obj, True, False) + wait_set.set_interest(writable_obj, False, True) + poller.events = [ + (31, select_mod.POLLHUP), + (32, select_mod.POLLERR), + ] + + readable, writable = wait_set.wait(timeout_ms=None) + wait_set.close() + + self.assertEqual([readable_obj], readable) + self.assertEqual([writable_obj], writable) + self.assertEqual([-1], poller.poll_timeouts) + + def test_micropython_wait_set_detaches_invalid_event(self): + poller = PersistentFakePoller() + select_mod = FakeSelectModule(poller) + kernel = MicroPythonKernel(modules={"select": select_mod}) + wait_set = kernel.create_io_wait_set() + io_obj = FakePollObject(41) + + wait_set.set_interest(io_obj, True, True) + poller.events = [(41, select_mod.POLLNVAL)] + readable, writable = wait_set.wait(timeout_ms=0) + poller.events = [] + second_readable, second_writable = wait_set.wait(timeout_ms=0) + wait_set.close() + + self.assertEqual([io_obj], readable) + self.assertEqual([io_obj], writable) + self.assertEqual([], second_readable) + self.assertEqual([], second_writable) + self.assertEqual([io_obj], poller.unregistrations) + + def test_micropython_poll_wait_set_works_with_cpython_poll(self): + kernel = MicroPythonKernel() + if kernel._poll_factory is None: + self.skipTest("host does not provide select.poll") + left, right = socket.socketpair() + wait_set = kernel.create_io_wait_set() + try: + wait_set.set_interest(left, True, False) + right.send(b"x") + + readable, writable = wait_set.wait(timeout_ms=100) + self.assertEqual([left], readable) + self.assertEqual([], writable) + finally: + wait_set.close() + left.close() + right.close() + + def test_micropython_kernel_without_poll_keeps_snapshot_fallback(self): + kernel = MicroPythonKernel(modules={"select": FakeSelectWithoutPoll()}) + + self.assertIsNone(kernel.create_io_wait_set()) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_runtime.py b/tests/test_runtime.py index a843162..288940e 100644 --- a/tests/test_runtime.py +++ b/tests/test_runtime.py @@ -1,9 +1,11 @@ +import os import sys +import socket import unittest sys.path.append("..") -from SmallPackage.Kernel import Kernel +from SmallPackage.Kernel import Kernel, Unix from SmallPackage.SmallOS import SmallOS from SmallPackage.SmallTask import SmallTask from SmallPackage.SmallErrors import TaskCancelledError @@ -64,6 +66,80 @@ def mark_writable(self, obj): self.writable.add(obj) +class StrictWaitKernel(FakeKernel): + def io_wait(self, readables, writables, timeout_ms=None): + watched = list(readables) + list(writables) + for obj in watched: + is_valid, exc = self.validate_io_wait_object(obj) + if not is_valid: + raise exc + return super().io_wait(readables, writables, timeout_ms=timeout_ms) + + +class CountingWaitSet: + def __init__(self, kernel): + self.kernel = kernel + self.interests = {} + self.interest_changes = [] + self.wait_calls = [] + self.closed = False + + def set_interest(self, obj, readable, writable): + interest = (bool(readable), bool(writable)) + previous = self.interests.get(obj, (False, False)) + if interest == previous: + return + self.interest_changes.append((obj, interest[0], interest[1])) + if interest == (False, False): + self.interests.pop(obj, None) + else: + self.interests[obj] = interest + + def wait(self, timeout_ms=None): + self.wait_calls.append(timeout_ms) + if timeout_ms is not None and timeout_ms > 0: + self.kernel.now += int(timeout_ms) + if timeout_ms == 0 and self.kernel.defer_ready_until_blocking: + return [], [] + + ready_read = [ + obj + for obj, interest in self.interests.items() + if interest[0] and obj in self.kernel.readable + ] + ready_write = [ + obj + for obj, interest in self.interests.items() + if interest[1] and obj in self.kernel.writable + ] + for obj in ready_read: + self.kernel.readable.discard(obj) + for obj in ready_write: + self.kernel.writable.discard(obj) + return ready_read, ready_write + + def close(self): + self.closed = True + self.interests.clear() + + +class PersistentFakeKernel(FakeKernel): + def __init__(self): + super().__init__() + self.wait_sets = [] + self.defer_ready_until_blocking = False + + def create_io_wait_set(self): + wait_set = CountingWaitSet(self) + self.wait_sets.append(wait_set) + return wait_set + + +class ClosedWaitObject: + def fileno(self): + return -1 + + class TestRuntime(unittest.TestCase): def build_os(self, *tasks): kernel = FakeKernel() @@ -73,6 +149,7 @@ def build_os(self, *tasks): return runtime, kernel def test_priority_order_is_preserved_before_and_after_sleep(self): + """Higher-priority tasks should run first both before and after sleeping.""" events = [] async def worker(task): @@ -92,6 +169,7 @@ async def worker(task): ) def test_join_all_returns_results_in_requested_order(self): + """join_all should preserve the caller-specified child ordering.""" async def child(task, value, delay): await task.sleep(delay) return value @@ -107,6 +185,7 @@ async def parent(task): self.assertEqual(["first", "second"], parent_task.result) def test_wait_signal_and_join_resume_once(self): + """A signal wait followed by join should resume exactly once per event.""" async def sender(task): await task.sleep(0.2) task.sendSignal(task.parent.pid, 7) @@ -124,6 +203,7 @@ async def parent(task): self.assertEqual((7, "sent"), parent_task.result) def test_killing_joined_child_raises_into_waiter(self): + """Cancelling a joined child should raise TaskCancelledError into the waiter.""" async def sleeper(task): await task.sleep(10) return "too late" @@ -148,6 +228,7 @@ async def parent(task): self.assertEqual("caught-cancel", parent_task.result) def test_wait_readable_resumes_on_kernel_io_event(self): + """A task waiting on readability should resume when the kernel marks it ready.""" io_obj = object() async def notifier(task, watched): @@ -165,7 +246,187 @@ async def waiter(task, watched): self.assertTrue(waiter_task.result) + def test_persistent_wait_set_uses_one_blocking_wait_on_idle_entry(self): + io_obj = object() + kernel = PersistentFakeKernel() + kernel.mark_readable(io_obj) + runtime = SmallOS().setKernel(kernel) + + async def waiter(task, watched): + return await task.wait_readable(watched) + + waiter_task = SmallTask(2, waiter, name="waiter", args=(io_obj,)) + runtime.fork(waiter_task) + runtime.startOS() + + wait_set = kernel.wait_sets[0] + self.assertIs(io_obj, waiter_task.result) + self.assertEqual([None], wait_set.wait_calls) + self.assertEqual( + [(io_obj, True, False), (io_obj, False, False)], + wait_set.interest_changes, + ) + self.assertTrue(wait_set.closed) + + def test_persistent_wait_set_modifies_combined_interest_only_on_changes(self): + io_obj = object() + kernel = PersistentFakeKernel() + kernel.defer_ready_until_blocking = True + kernel.mark_readable(io_obj) + kernel.mark_writable(io_obj) + runtime = SmallOS().setKernel(kernel) + + async def read_waiter(task, watched): + return await task.wait_readable(watched) + + async def write_waiter(task, watched): + return await task.wait_writable(watched) + + reader = SmallTask(2, read_waiter, name="reader", args=(io_obj,)) + writer = SmallTask(3, write_waiter, name="writer", args=(io_obj,)) + runtime.fork([reader, writer]) + runtime.startOS() + + wait_set = kernel.wait_sets[0] + self.assertIs(io_obj, reader.result) + self.assertIs(io_obj, writer.result) + self.assertEqual( + [ + (io_obj, True, False), + (io_obj, True, True), + (io_obj, False, False), + ], + wait_set.interest_changes, + ) + self.assertEqual([0, None], wait_set.wait_calls) + + def test_persistent_wait_set_is_rebuilt_when_watcher_runtime_restarts(self): + io_obj = object() + kernel = PersistentFakeKernel() + runtime = SmallOS().setKernel(kernel) + + async def watcher(task, watched): + return await task.wait_readable(watched) + + watcher_task = SmallTask( + 2, + watcher, + name="watcher", + args=(io_obj,), + isWatcher=True, + ) + runtime.fork(watcher_task) + runtime.startOS() + + self.assertEqual(1, len(kernel.wait_sets)) + self.assertTrue(kernel.wait_sets[0].closed) + self.assertIn(io_obj, runtime.ioReadWaiters) + + kernel.mark_readable(io_obj) + runtime.setEternalWatchers(True) + runtime.startOS() + + self.assertEqual(2, len(kernel.wait_sets)) + self.assertTrue(kernel.wait_sets[1].closed) + self.assertIs(io_obj, watcher_task.result) + self.assertNotIn(io_obj, runtime.ioReadWaiters) + + def test_cancelling_persistent_io_waiter_unregisters_last_interest(self): + io_obj = object() + kernel = PersistentFakeKernel() + runtime = SmallOS().setKernel(kernel) + + async def io_waiter(task, watched): + await task.wait_readable(watched) + + async def killer(task, target): + await task.sleep(0.1) + target.kill() + + async def parent(task, watched): + waiter = task.spawn(io_waiter, priority=2, name="waiter", args=(watched,)) + task.spawn(killer, priority=1, name="killer", args=(waiter,)) + await task.sleep(0.2) + return watched not in task.OS.ioReadWaiters + + parent_task = SmallTask(3, parent, name="parent", args=(io_obj,)) + runtime.fork(parent_task) + runtime.startOS() + + self.assertTrue(parent_task.result) + self.assertEqual( + [(io_obj, True, False), (io_obj, False, False)], + kernel.wait_sets[0].interest_changes, + ) + + def test_unix_invalid_registration_fails_await_not_scheduler(self): + left, right = socket.socketpair() + left.close() + runtime = SmallOS().setKernel(Unix()) + + async def waiter(task, watched): + try: + await task.wait_readable(watched) + except ValueError as exc: + return "invalid file descriptor" in str(exc) + return False + + waiter_task = SmallTask(2, waiter, name="waiter", args=(left,)) + runtime.fork(waiter_task) + try: + runtime.startOS() + finally: + right.close() + + self.assertTrue(waiter_task.result) + + def test_unix_closed_integer_descriptor_fails_await_not_scheduler(self): + read_fd, write_fd = os.pipe() + os.close(read_fd) + runtime = SmallOS().setKernel(Unix()) + + async def waiter(task, watched): + try: + await task.wait_readable(watched) + except OSError: + return True + return False + + waiter_task = SmallTask(2, waiter, name="waiter", args=(read_fd,)) + runtime.fork(waiter_task) + try: + runtime.startOS() + finally: + os.close(write_fd) + + self.assertTrue(waiter_task.result) + + def test_unix_runtime_uses_persistent_selector_for_socket_wait(self): + left, right = socket.socketpair() + left.setblocking(False) + right.setblocking(False) + runtime = SmallOS().setKernel(Unix()) + + async def reader(task, sock): + ready = await task.wait_readable(sock) + return ready.recv(1) + + async def writer(task, sock): + sock.send(b"x") + + reader_task = SmallTask(1, reader, name="reader", args=(left,)) + writer_task = SmallTask(2, writer, name="writer", args=(right,)) + runtime.fork([reader_task, writer_task]) + try: + runtime.startOS() + finally: + left.close() + right.close() + + self.assertEqual(b"x", reader_task.result) + 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() async def io_waiter(task, watched): @@ -185,9 +446,349 @@ async def parent(task, watched): parent_task = SmallTask(2, parent, name="parent", args=(io_obj,)) self.build_os(parent_task) + self.assertTrue(parent_task.result) + + def test_cancelling_signal_waiter_clears_wait_metadata(self): + """Cancelling a signal waiter should clear its signal-specific blocked metadata.""" + async def signal_waiter(task): + await task.wait_signal(9) + return "unexpected" + + async def killer(task, target): + await task.sleep(0.2) + target.kill() + return "killed" + + async def parent(task): + waiter = task.spawn(signal_waiter, priority=3, name="signal_waiter") + task.spawn(killer, priority=1, name="killer", args=(waiter,)) + await task.sleep(0.5) + return ( + waiter._waiting_signal is None + and waiter._blocked_reason is None + and isinstance(waiter.exception, TaskCancelledError) + ) + + parent_task = SmallTask(2, parent, name="parent") + self.build_os(parent_task) self.assertTrue(parent_task.result) + def test_cancelling_join_waiter_unregisters_child(self): + """Cancelling a join waiter should unregister it from the joined child.""" + async def sleeper(task): + await task.sleep(1) + return "done" + + async def join_waiter(task, target): + await task.join(target) + return "unexpected" + + async def killer(task, target): + await task.sleep(0.2) + target.kill() + return "killed" + + async def parent(task): + child = task.spawn(sleeper, priority=4, name="sleeper") + waiter = task.spawn(join_waiter, priority=3, name="join_waiter", args=(child,)) + task.spawn(killer, priority=1, name="killer", args=(waiter,)) + await task.sleep(0.5) + return ( + not child._join_waiters + and waiter._join_target is None + and waiter._blocked_reason is None + and isinstance(waiter.exception, TaskCancelledError) + ) + + parent_task = SmallTask(2, parent, name="parent") + self.build_os(parent_task) + + self.assertTrue(parent_task.result) + + def test_cancelling_join_all_waiter_unregisters_children(self): + """Cancelling a join_all waiter should unregister it from every joined child.""" + async def sleeper(task): + await task.sleep(1) + return "done" + + async def join_all_waiter(task, targets): + await task.join_all(targets) + return "unexpected" + + async def killer(task, target): + await task.sleep(0.2) + target.kill() + return "killed" + + async def parent(task): + first = task.spawn(sleeper, priority=4, name="first") + second = task.spawn(sleeper, priority=4, name="second") + waiter = task.spawn( + join_all_waiter, + priority=3, + name="join_all_waiter", + args=([first, second],), + ) + task.spawn(killer, priority=1, name="killer", args=(waiter,)) + await task.sleep(0.5) + return ( + not first._join_waiters + and not second._join_waiters + and waiter._join_targets is None + and waiter._join_pending == set() + and waiter._blocked_reason is None + and isinstance(waiter.exception, TaskCancelledError) + ) + + parent_task = SmallTask(2, parent, name="parent") + self.build_os(parent_task) + + self.assertTrue(parent_task.result) + + def test_resuming_task_clears_previous_wait_metadata_before_next_wait(self): + """Resuming from one wait should clear old wait metadata before the next wait.""" + io_obj = object() + + async def waiter_task(task, watched): + await task.wait_signal(7) + ready_obj = await task.wait_readable(watched) + return ready_obj is watched + + async def notifier(task, target, watched): + await task.sleep(0.2) + task.sendSignal(target.getID(), 7) + await task.sleep(0.3) + task.OS.kernel.mark_readable(watched) + return "notified" + + async def inspector(task, target, watched): + await task.sleep(0.3) + return ( + target._waiting_signal is None + and target._io_wait_obj is watched + and target._io_wait_mode == "read" + and target._blocked_reason == "wait_readable" + ) + + async def parent(task, watched): + waiter = task.spawn(waiter_task, priority=3, name="waiter", args=(watched,)) + inspector_task = task.spawn( + inspector, + priority=1, + name="inspector", + args=(waiter, watched), + ) + task.spawn(notifier, priority=1, name="notifier", args=(waiter, watched)) + inspect_ok = await task.join(inspector_task) + waiter_ok = await task.join(waiter) + return inspect_ok and waiter_ok + + parent_task = SmallTask(2, parent, name="parent", args=(io_obj,)) + self.build_os(parent_task) + + self.assertTrue(parent_task.result) + + def test_error_handler_receives_uncaught_task_failure(self): + """The runtime error handler should receive uncaught task failures once.""" + events = [] + kernel = FakeKernel() + runtime = SmallOS().setKernel(kernel).setErrorHandler(events.append) + + async def boom(task): + raise RuntimeError("boom") + + root = SmallTask(2, boom, name="boom") + runtime.fork([root]) + runtime.startOS() + + self.assertEqual(1, len(events)) + event = events[0] + self.assertEqual(root.getID(), event["task_id"]) + self.assertEqual("boom", event["task_name"]) + self.assertIsNone(event["parent_id"]) + self.assertEqual("RuntimeError", event["exception_type"]) + self.assertEqual("RuntimeError('boom')", event["exception_repr"]) + self.assertFalse(event["is_cancelled"]) + self.assertIsNone(event["blocked_reason"]) + self.assertIsNone(event["waiting_signal"]) + self.assertIsNone(event["io_wait_mode"]) + self.assertIsNone(event["join_target_id"]) + self.assertEqual([], event["join_pending_ids"]) + self.assertIsInstance(event["exception"], RuntimeError) + self.assertIn("RuntimeError: boom", event["traceback_text"]) + + def test_error_handler_ignores_successful_completion(self): + """Successful task completion should not trigger the runtime error handler.""" + events = [] + kernel = FakeKernel() + runtime = SmallOS().setKernel(kernel).setErrorHandler(events.append) + + async def worker(task): + await task.sleep(0.1) + return "ok" + + root = SmallTask(2, worker, name="worker") + runtime.fork([root]) + runtime.startOS() + + self.assertEqual("ok", root.result) + self.assertEqual([], events) + + def test_error_handler_ignores_cancelled_tasks_by_default(self): + """TaskCancelledError should be ignored by the handler unless explicitly included.""" + events = [] + kernel = FakeKernel() + runtime = SmallOS().setKernel(kernel).setErrorHandler(events.append) + + async def signal_waiter(task): + await task.wait_signal(9) + return "unexpected" + + async def killer(task, target): + await task.sleep(0.2) + target.kill() + return "killed" + + async def parent(task): + waiter = task.spawn(signal_waiter, priority=3, name="signal_waiter") + task.spawn(killer, priority=1, name="killer", args=(waiter,)) + await task.sleep(0.5) + return isinstance(waiter.exception, TaskCancelledError) + + parent_task = SmallTask(2, parent, name="parent") + runtime.fork([parent_task]) + runtime.startOS() + + self.assertTrue(parent_task.result) + self.assertEqual([], events) + + def test_error_handler_can_include_cancelled_tasks_with_pre_finalize_snapshot(self): + """The handler can opt into cancellations and still see pre-finalize wait context.""" + events = [] + captured = {} + kernel = FakeKernel() + runtime = SmallOS().setKernel(kernel).setErrorHandler(events.append, include_cancelled=True) + + async def signal_waiter(task): + await task.wait_signal(9) + return "unexpected" + + async def killer(task, target): + await task.sleep(0.2) + target.kill() + return "killed" + + async def parent(task): + waiter = task.spawn(signal_waiter, priority=3, name="signal_waiter") + captured["waiter"] = waiter + task.spawn(killer, priority=1, name="killer", args=(waiter,)) + await task.sleep(0.5) + return "done" + + parent_task = SmallTask(2, parent, name="parent") + runtime.fork([parent_task]) + runtime.startOS() + + waiter = captured["waiter"] + self.assertEqual("done", parent_task.result) + self.assertEqual(1, len(events)) + event = events[0] + self.assertEqual(waiter.getID(), event["task_id"]) + self.assertTrue(event["is_cancelled"]) + self.assertEqual("TaskCancelledError", event["exception_type"]) + self.assertEqual("signal", event["blocked_reason"]) + self.assertEqual(9, event["waiting_signal"]) + self.assertIsNone(event["io_wait_mode"]) + self.assertIsNone(waiter._blocked_reason) + self.assertIsNone(waiter._waiting_signal) + + def test_error_handler_failure_is_reported_without_crashing_runtime(self): + """Failures inside the runtime error handler should be downgraded to diagnostics.""" + calls = [] + kernel = FakeKernel() + + def broken_handler(event): + calls.append(event["task_id"]) + raise RuntimeError("handler boom") + + runtime = SmallOS().setKernel(kernel).setErrorHandler(broken_handler) + + async def boom(task): + raise ValueError("task boom") + + root = SmallTask(2, boom, name="boom") + runtime.fork([root]) + runtime.startOS() + + self.assertEqual([root.getID()], calls) + self.assertIsInstance(root.exception, ValueError) + self.assertTrue( + any("smallOS error handler failed:" in message for message in kernel.output) + ) + + def test_invalid_io_wait_object_resumes_waiter_with_error(self): + """An invalid I/O wait object should resume the waiter with a clean exception.""" + closed_obj = ClosedWaitObject() + kernel = StrictWaitKernel() + runtime = SmallOS().setKernel(kernel) + + async def waiter(task, watched): + try: + await task.wait_readable(watched) + except ValueError as exc: + return str(exc) + return "unexpected" + + waiter_task = SmallTask(2, waiter, name="waiter", args=(closed_obj,)) + runtime.fork([waiter_task]) + runtime.startOS() + + self.assertIn("invalid file descriptor (-1)", waiter_task.result) + self.assertNotIn(closed_obj, runtime.ioReadWaiters) + + def test_uncaught_invalid_io_wait_error_clears_wait_state(self): + """An uncaught invalid-I/O failure should still clear runtime wait bookkeeping.""" + closed_obj = ClosedWaitObject() + kernel = StrictWaitKernel() + runtime = SmallOS().setKernel(kernel) + + async def waiter(task, watched): + await task.wait_readable(watched) + return "unexpected" + + waiter_task = SmallTask(2, waiter, name="waiter", args=(closed_obj,)) + runtime.fork([waiter_task]) + runtime.startOS() + + self.assertIsInstance(waiter_task.exception, ValueError) + self.assertIsNone(waiter_task._io_wait_obj) + self.assertIsNone(waiter_task._io_wait_mode) + self.assertIsNone(waiter_task._blocked_reason) + self.assertNotIn(closed_obj, runtime.ioReadWaiters) + + def test_error_handler_receives_invalid_io_failure(self): + """Invalid-I/O task failures should also be delivered through the error handler.""" + events = [] + closed_obj = ClosedWaitObject() + kernel = StrictWaitKernel() + runtime = SmallOS().setKernel(kernel).setErrorHandler(events.append) + + async def waiter(task, watched): + await task.wait_readable(watched) + return "unexpected" + + waiter_task = SmallTask(2, waiter, name="waiter", args=(closed_obj,)) + runtime.fork([waiter_task]) + runtime.startOS() + + self.assertEqual(1, len(events)) + event = events[0] + self.assertEqual(waiter_task.getID(), event["task_id"]) + self.assertEqual("ValueError", event["exception_type"]) + self.assertIn("invalid file descriptor (-1)", event["exception_repr"]) + self.assertIn("invalid file descriptor (-1)", event["traceback_text"]) + if __name__ == "__main__": unittest.main()