diff --git a/learning_loop_node/tests/unit/test_subprocess.py b/learning_loop_node/tests/unit/test_subprocess.py index b197cc9b..1b2c6418 100644 --- a/learning_loop_node/tests/unit/test_subprocess.py +++ b/learning_loop_node/tests/unit/test_subprocess.py @@ -1,5 +1,16 @@ +import asyncio +import multiprocessing +import os +import queue +import signal +import sys +from collections.abc import Callable, Iterator +from typing import Any + import pytest +from ...trainer import subprocess as iterator_module +from ...trainer.exceptions import UnexpectedWorkerExitError from ...trainer.subprocess import iterator_cpu_bound @@ -51,3 +62,80 @@ async def test_leaving_early_does_not_leave_the_process_running(): if item == 2: break # reaching here without hanging is the test + + +@pytest.mark.parametrize(('mode', 'exit_code'), [ + ('system_exit', 1), + ('system_exit_zero', 0), + ('native_exit', 7), + ('kill', -signal.SIGKILL), +]) +async def test_iterator_rejects_exit_without_completion(mode: str, exit_code: int) -> None: + with pytest.raises(UnexpectedWorkerExitError, match=rf'exited with code {exit_code} .*IteratorDone'): + await asyncio.wait_for(_collect_worker(_exit_without_completion, mode), timeout=10) + + +@pytest.mark.parametrize('error', [SystemExit(1), ValueError('worker failed')]) +async def test_iterator_reports_failure_after_progress(error: BaseException) -> None: + seen = [] + + async def consume() -> None: + async with iterator_cpu_bound(_fail_after_progress, error) as results: + async for item in results: + seen.append(item) + + expected_error = UnexpectedWorkerExitError if isinstance(error, SystemExit) else ValueError + with pytest.raises(expected_error): + await asyncio.wait_for(consume(), timeout=10) + assert seen == [0] + + +async def test_iterator_accepts_completion_arriving_after_queue_timeout(monkeypatch: pytest.MonkeyPatch) -> None: + context = multiprocessing.get_context('spawn') + make_process = context.Process + processes = [] + + def record_process(*args: Any, **kwargs: Any) -> Any: + process = make_process(*args, **kwargs) + processes.append(process) + return process + + to_thread = asyncio.to_thread + first_poll = True + + async def delayed_timeout(func: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + nonlocal first_poll + if first_poll: + first_poll = False + await to_thread(processes[0].join, 5) + assert processes[0].exitcode == 0 + raise queue.Empty + return await to_thread(func, *args, **kwargs) + + monkeypatch.setattr(context, 'Process', record_process) + monkeypatch.setattr(iterator_module.asyncio, 'to_thread', delayed_timeout) + + assert await asyncio.wait_for(_collect_worker(counting, 0), timeout=10) == [] + + +async def _collect_worker(it: Callable[..., Iterator[int]], *args: Any) -> list[int]: + async with iterator_cpu_bound(it, *args) as results: + return [item async for item in results] + + +def _exit_without_completion(mode: str) -> Iterator[int]: + yield from () + if mode == 'system_exit': + sys.exit(1) + if mode == 'system_exit_zero': + sys.exit(0) + if mode == 'native_exit': + os._exit(7) + if mode == 'kill': + os.kill(os.getpid(), signal.SIGKILL) + raise AssertionError(f'Unexpected exit mode: {mode}') + + +def _fail_after_progress(error: BaseException) -> Iterator[int]: + yield 0 + raise error diff --git a/learning_loop_node/trainer/exceptions.py b/learning_loop_node/trainer/exceptions.py index b77548cd..6eec5fa8 100644 --- a/learning_loop_node/trainer/exceptions.py +++ b/learning_loop_node/trainer/exceptions.py @@ -9,6 +9,10 @@ class InsufficientMemoryError(RuntimeError): """Raised when not even the smallest unit of work fits in memory.""" +class UnexpectedWorkerExitError(RuntimeError): + """Raised when an iterator worker exits without reporting completion.""" + + class NodeNeedsRestartError(Exception): """ NodeNeedsRestartError is raised when the node needs to be restarted. diff --git a/learning_loop_node/trainer/subprocess.py b/learning_loop_node/trainer/subprocess.py index 118fbf5d..50a9c7b9 100644 --- a/learning_loop_node/trainer/subprocess.py +++ b/learning_loop_node/trainer/subprocess.py @@ -16,6 +16,8 @@ from multiprocessing.queues import Queue as MPQueue from typing import Any, ParamSpec, TypeVar +from .exceptions import UnexpectedWorkerExitError + logger = logging.getLogger(__name__) T = TypeVar('T') @@ -55,9 +57,15 @@ async def _iterator_cpu_bound_inner( try: item = await asyncio.to_thread(state_queue.get, True, 0.5) except queue.Empty: - if not process.is_alive(): - break - continue + if process.is_alive(): + continue + # Completion may have arrived between the timeout and the exit check. + try: + item = await asyncio.to_thread(state_queue.get_nowait) + except queue.Empty as e: + raise UnexpectedWorkerExitError( + f'{process.name} exited with code {process.exitcode} without sending IteratorDone' + ) from e match item: case IteratorDone(): break