From 353c25add5f2ef58316c3624653a7c9e74f8e3f1 Mon Sep 17 00:00:00 2001 From: Itay Dar <118370953+ItayTheDar@users.noreply.github.com> Date: Sat, 25 Jul 2026 22:37:46 +0300 Subject: [PATCH] feat: add supervised background workers --- docs/background_workers.md | 221 +++++++++++++ docs/lifespan_tasks.md | 51 ++- mkdocs.yml | 1 + nest/common/__init__.py | 7 + nest/common/background_worker.py | 132 ++++++++ nest/core/__init__.py | 13 + nest/core/pynest_application.py | 47 ++- nest/core/pynest_container.py | 31 +- nest/core/pynest_factory.py | 13 +- nest/core/worker_factory.py | 204 ++++++++++++ nest/core/worker_host.py | 340 ++++++++++++++++++++ tests/test_common/test_background_worker.py | 167 ++++++++++ tests/test_core/test_worker_host.py | 325 +++++++++++++++++++ tests/test_core/test_worker_integration.py | 238 ++++++++++++++ 14 files changed, 1746 insertions(+), 44 deletions(-) create mode 100644 docs/background_workers.md create mode 100644 nest/common/background_worker.py create mode 100644 nest/core/worker_factory.py create mode 100644 nest/core/worker_host.py create mode 100644 tests/test_common/test_background_worker.py create mode 100644 tests/test_core/test_worker_host.py create mode 100644 tests/test_core/test_worker_integration.py diff --git a/docs/background_workers.md b/docs/background_workers.md new file mode 100644 index 0000000..7d4005b --- /dev/null +++ b/docs/background_workers.md @@ -0,0 +1,221 @@ +# Background workers + +PyNest can run dependency-injected work alongside an HTTP application or in a +dedicated process without an HTTP server. The framework starts worker tasks on +the runtime event loop, supervises failures, and stops them before application +shutdown hooks dispose their dependencies. + +## Long-running workers + +Create an injectable provider that extends `BackgroundWorker`: + +```python +import asyncio + +from nest.core import BackgroundWorker, Injectable, Module + + +@Injectable +class EmailWorker(BackgroundWorker): + name = "email" + + def __init__(self, queue: EmailQueue): + self.queue = queue + + async def run(self) -> None: + while not self.stopping.is_set(): + try: + message = await asyncio.wait_for( + self.queue.receive(), + timeout=1, + ) + except asyncio.TimeoutError: + continue + + await self.queue.deliver(message) + + +@Module(providers=[EmailQueue, EmailWorker]) +class AppModule: + pass +``` + +Register the worker in `providers` like any other service. Constructor +dependencies are resolved by the PyNest container, and the host retains the +resolved instance for the application lifespan. Singleton scope is recommended +for worker providers. + +Do not start tasks from a worker constructor or lifecycle hook. PyNest calls +`run()` from FastAPI's lifespan event loop, which keeps asyncio resources on the +same loop that owns the application. + +### Cooperative shutdown + +`self.stopping` is set when shutdown begins. Long-running loops should inspect +it, and idle delays should use `self.sleep()`: + +```python +async def run(self) -> None: + while not self.stopping.is_set(): + await self.flush_batch() + if not await self.sleep(5): + return +``` + +`sleep(seconds)` returns `True` when the delay elapsed and `False` when shutdown +interrupted it. PyNest first requests cooperative shutdown and calls +`on_stop()`. If the worker is still running after the grace timeout, its task is +cancelled and drained. + +`on_start()` and `on_stop()` are optional. They may be synchronous or +asynchronous, but `run()` must be asynchronous. + +## Recurring interval jobs + +`IntervalWorker` is a fixed-delay scheduler built on the same supervisor: + +```python +from nest.core import Injectable, IntervalWorker + + +@Injectable +class CleanupWorker(IntervalWorker): + name = "expired-session-cleanup" + interval = 300 + run_immediately = True + + def __init__(self, sessions: SessionRepository): + self.sessions = sessions + + async def execute(self) -> None: + await self.sessions.delete_expired() +``` + +The delay starts after `execute()` finishes, so occurrences never overlap +within one process. `interval` must be greater than zero. By default, the first +occurrence waits for one interval; set `run_immediately = True` to run once at +startup. + +Calendar and cron scheduling are intentionally not implemented in core. +Production cron scheduling needs explicit policies for time zones, daylight +saving changes, missed executions, persistence, and distributed locking. Use a +dedicated scheduler in a worker provider when those semantics are required. + +## Restart policies + +Workers default to restarting after an exception: + +```python +from nest.core import RestartPolicy + + +class QueueWorker(BackgroundWorker): + restart = RestartPolicy.ON_FAILURE + restart_backoff = 1 + restart_backoff_max = 30 +``` + +The available policies are: + +| Policy | Exception | Normal return | +| --- | --- | --- | +| `RestartPolicy.NONE` | Stop as failed | Stop as completed | +| `RestartPolicy.ON_FAILURE` | Restart | Stop as completed | +| `RestartPolicy.ALWAYS` | Restart | Restart | + +Restarts use capped exponential backoff. Cancellation and application shutdown +never trigger a restart. + +An unhandled `IntervalWorker.execute()` exception follows the same policy as +any other worker failure. + +## HTTP applications + +No extra startup code is needed: + +```python +from nest.core import PyNestFactory + +app = PyNestFactory.create( + AppModule, + worker_grace_timeout=15, + title="API and workers", +) +``` + +Workers start when the ASGI lifespan starts, not when +`PyNestFactory.create()` returns. They stop before PyNest runs provider and +module shutdown hooks. + +Inspect worker state for a health or administration endpoint: + +```python +worker_status = app.get_worker_host().status() +``` + +The result is JSON-ready: + +```json +[ + { + "name": "email", + "state": "running", + "restarts": 1, + "last_error": "ConnectionError: broker unavailable" + } +] +``` + +States are `starting`, `running`, `backing_off`, `stopping`, `stopped`, +`completed`, and `failed`. + +## Standalone worker processes + +Use `run_workers()` for a process that does not serve HTTP: + +```python +from nest.core import run_workers + +from src.app_module import AppModule + + +if __name__ == "__main__": + run_workers(AppModule, grace_timeout=15) +``` + +The standalone runner: + +1. Builds the dependency container. +2. Runs bootstrap lifecycle hooks on its event loop. +3. Starts all registered workers on that same loop. +4. Handles `SIGTERM` and `SIGINT`. +5. Stops workers before running shutdown hooks. + +For an existing async entrypoint, use the application object instead: + +```python +from nest.core import WorkerAppFactory + + +async def main() -> None: + app = WorkerAppFactory.create(AppModule) + await app.run() +``` + +Do not call `run_workers()` from a running event loop; it raises a clear error +instead of nesting `asyncio.run()`. + +## Deployment behavior + +Every application process runs its own instance of every registered worker. If +Uvicorn starts four processes, an HTTP-integrated interval worker runs four +times. + +Use one of these patterns when work must run only once: + +- Deploy `run_workers()` as a separate single-replica service. +- Partition queue consumption so the broker coordinates consumers. +- Protect scheduled work with a distributed lock. +- Use a scheduler that persists and coordinates jobs across processes. + +Worker status is local to one process and resets when that process restarts. diff --git a/docs/lifespan_tasks.md b/docs/lifespan_tasks.md index f71a6aa..f19cadd 100644 --- a/docs/lifespan_tasks.md +++ b/docs/lifespan_tasks.md @@ -1,42 +1,37 @@ -# Lifaspan tasks in PyNest +# Lifespan tasks in PyNest -## Introduction +Long-running coroutines should use PyNest's +[background worker](background_workers.md) support. Awaiting an infinite +coroutine directly from a FastAPI startup handler prevents startup from +completing and does not provide supervised failure or bounded shutdown. -Lifespan tasks - coroutines, which run while app is working. - -## Defining a lifespan task -As example of lifespan task will use coroutine, which print time every hour. In real user cases can be everything else. +For a recurring task, extend `IntervalWorker`: ```python -import asyncio from datetime import datetime -async def print_current_time(): - while True: +from nest.core import Injectable, IntervalWorker, Module, PyNestFactory + + +@Injectable +class ClockWorker(IntervalWorker): + interval = 3600 + run_immediately = True + + async def execute(self) -> None: current_time = datetime.now().strftime("%H:%M:%S") print(f"Current time: {current_time}") - await asyncio.sleep(3600) -``` - -## Implement a lifespan task -In `app_module.py` we can define a startup handler, and run lifespan inside it -```python -from nest.core import PyNestFactory -app = PyNestFactory.create( - AppModule, - description="This is my PyNest app with lifespan task", - title="My App", - version="1.0.0", - debug=True, -) +@Module(providers=[ClockWorker]) +class AppModule: + pass -http_server = app.get_server() -@http_server.on_event("startup") -async def startup(): - await print_current_time() +app = PyNestFactory.create(AppModule) ``` -Now `print_current_time` will work in lifespan after startup. +PyNest starts the worker inside the ASGI lifespan and stops it when the +application shuts down. See [Background workers](background_workers.md) for +long-running consumers, restart policies, status inspection, and standalone +worker processes. diff --git a/mkdocs.yml b/mkdocs.yml index 202ee6d..637a96b 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -60,6 +60,7 @@ nav: - Guards: guards.md - Exception Filters: exception_filters.md - WebSockets: websockets.md + - Background Workers: background_workers.md - Dependency Injection: dependency_injection.md - Deployment: - Docker: docker.md diff --git a/nest/common/__init__.py b/nest/common/__init__.py index 6da2eb3..d49b265 100644 --- a/nest/common/__init__.py +++ b/nest/common/__init__.py @@ -31,3 +31,10 @@ OnModuleDestroy, OnModuleInit, ) +from nest.common.background_worker import ( + BackgroundWorker, + IntervalWorker, + RestartPolicy, + WorkerState, + WorkerStatus, +) diff --git a/nest/common/background_worker.py b/nest/common/background_worker.py new file mode 100644 index 0000000..d0db0f8 --- /dev/null +++ b/nest/common/background_worker.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +import asyncio +from abc import ABC, abstractmethod +from dataclasses import dataclass +from enum import Enum +from typing import Optional + + +class RestartPolicy(str, Enum): + """Controls when a worker is restarted after ``run`` exits.""" + + NONE = "none" + ON_FAILURE = "on_failure" + ALWAYS = "always" + + +class WorkerState(str, Enum): + """Observable state of a managed background worker.""" + + STARTING = "starting" + RUNNING = "running" + BACKING_OFF = "backing_off" + STOPPING = "stopping" + STOPPED = "stopped" + COMPLETED = "completed" + FAILED = "failed" + + +@dataclass(frozen=True) +class WorkerStatus: + """Immutable snapshot returned by the worker host.""" + + name: str + state: WorkerState + restarts: int = 0 + last_error: Optional[str] = None + + def as_dict(self) -> dict[str, object]: + """Return a JSON-serializable status representation.""" + return { + "name": self.name, + "state": self.state.value, + "restarts": self.restarts, + "last_error": self.last_error, + } + + +class BackgroundWorker(ABC): + """ + Base class for dependency-injected, long-running application work. + + Worker subclasses do not need to call ``super().__init__()``. PyNest owns + the stop event and recreates it each time the worker host starts. + """ + + name: str = "" + restart: RestartPolicy = RestartPolicy.ON_FAILURE + restart_backoff: float = 1.0 + restart_backoff_max: float = 30.0 + + @abstractmethod + async def run(self) -> None: + """Run until the work completes or cooperative shutdown is requested.""" + raise NotImplementedError + + async def on_start(self) -> None: + """Run once immediately before the worker is supervised.""" + + async def on_stop(self) -> None: + """Run once after cooperative shutdown has been requested.""" + + @property + def stopping(self) -> asyncio.Event: + """Event set by PyNest when the worker should stop.""" + event = getattr(self, "_pynest_stopping_event", None) + if event is None: + event = asyncio.Event() + self._pynest_stopping_event = event + return event + + async def sleep(self, seconds: float) -> bool: + """ + Sleep without delaying application shutdown. + + Returns ``True`` when the delay elapsed, or ``False`` when shutdown + interrupted it. + """ + if seconds < 0: + raise ValueError("seconds must be greater than or equal to zero") + if self.stopping.is_set(): + return False + try: + await asyncio.wait_for(self.stopping.wait(), timeout=seconds) + except asyncio.TimeoutError: + return True + return False + + def _prepare_for_start(self) -> None: + self._pynest_stopping_event = asyncio.Event() + + def _request_stop(self) -> None: + self.stopping.set() + + +class IntervalWorker(BackgroundWorker): + """ + A fixed-delay recurring worker. + + The interval starts after each ``execute`` call completes, so executions + never overlap within one process. + """ + + interval: float = 60.0 + run_immediately: bool = False + + @abstractmethod + async def execute(self) -> None: + """Execute one occurrence of the recurring job.""" + raise NotImplementedError + + async def run(self) -> None: + if self.interval <= 0: + raise ValueError("interval must be greater than zero") + + if not self.run_immediately and not await self.sleep(self.interval): + return + + while not self.stopping.is_set(): + await self.execute() + if not await self.sleep(self.interval): + return diff --git a/nest/core/__init__.py b/nest/core/__init__.py index a5cf3a5..826b5c2 100644 --- a/nest/core/__init__.py +++ b/nest/core/__init__.py @@ -13,6 +13,13 @@ createParamDecorator, ) from nest.common.provider import InjectionToken, Scope +from nest.common.background_worker import ( + BackgroundWorker, + IntervalWorker, + RestartPolicy, + WorkerState, + WorkerStatus, +) from nest.core.decorators import ( Catch, Controller, @@ -30,3 +37,9 @@ from nest.core.pynest_application import PyNestApp from nest.core.pynest_container import PyNestContainer from nest.core.pynest_factory import PyNestFactory +from nest.core.worker_factory import ( + WorkerApplication, + WorkerAppFactory, + run_workers, +) +from nest.core.worker_host import WorkerHost diff --git a/nest/core/pynest_application.py b/nest/core/pynest_application.py index 45f2295..7e6d41d 100644 --- a/nest/core/pynest_application.py +++ b/nest/core/pynest_application.py @@ -11,6 +11,7 @@ from nest.common.route_resolver import RoutesResolver from nest.core.pynest_container import PyNestContainer +from nest.core.worker_host import WorkerHost class PyNestApp: @@ -18,11 +19,20 @@ class PyNestApp: Main PyNest application. Wraps a built container and a FastAPI HTTP server. """ - def __init__(self, container: PyNestContainer, http_server: FastAPI) -> None: + def __init__( + self, + container: PyNestContainer, + http_server: FastAPI, + *, + worker_grace_timeout: float = 10.0, + ) -> None: self.container = container self.http_server = http_server + self._worker_host = WorkerHost.discover( + container, grace_timeout=worker_grace_timeout + ) self._closed = False - self._closing = False + self._close_task: Optional[asyncio.Task[None]] = None self._install_lifespan_shutdown() routes_resolver = RoutesResolver(self.container, self.http_server) routes_resolver.register_routes() @@ -34,6 +44,10 @@ def get_http_server(self) -> FastAPI: """Alias for get_server() — kept for backward compatibility.""" return self.http_server + def get_worker_host(self) -> WorkerHost: + """Return the application's background worker supervisor.""" + return self._worker_host + def use(self, middleware: type, **options: Any) -> "PyNestApp": """Add ASGI middleware to the FastAPI server.""" self.http_server.add_middleware(middleware, **options) @@ -54,15 +68,21 @@ def enable_shutdown_hooks( async def close(self, signal: Optional[str] = None) -> None: """Run graceful application shutdown lifecycle hooks once.""" - if self._closed or self._closing: + if self._close_task is not None: + await asyncio.shield(self._close_task) + return + if self._closed: return - self._closing = True - try: - await self.container.shutdown_lifecycle(signal) - self._closed = True - finally: - self._closing = False + self._close_task = asyncio.create_task( + self._close(signal), name="pynest-application-close" + ) + await asyncio.shield(self._close_task) + + async def _close(self, signal: Optional[str]) -> None: + await self._worker_host.stop() + await self.container.shutdown_lifecycle(signal) + self._closed = True def use_global_filters(self, *filters) -> "PyNestApp": """Register one or more exception filters that apply to every route. @@ -132,10 +152,11 @@ def _install_lifespan_shutdown(self) -> None: @asynccontextmanager async def lifespan_context(app: FastAPI): - async with original_lifespan_context(app) as state: - try: + try: + async with original_lifespan_context(app) as state: + await self._worker_host.start() yield state - finally: - await self.close() + finally: + await self.close() self.http_server.router.lifespan_context = lifespan_context diff --git a/nest/core/pynest_container.py b/nest/core/pynest_container.py index 120dc2c..f949a4c 100644 --- a/nest/core/pynest_container.py +++ b/nest/core/pynest_container.py @@ -2,7 +2,7 @@ import inspect import logging -from typing import Any, Dict, List, Optional, Type, Union +from typing import Any, Dict, List, Optional, Type, TypeVar, Union from nest.common.exceptions import CircularDependencyException from nest.common.interfaces import ( @@ -26,6 +26,8 @@ "on_application_shutdown", ) +T = TypeVar("T") + class ModuleRef: """Internal container representation of a registered module.""" @@ -117,6 +119,33 @@ def get_controller_instance(self, controller_class: Type) -> Any: """Get a controller instance with all its service dependencies injected.""" return self.get(controller_class) + def get_provider_instances(self) -> List[Any]: + """Return every registered provider instance once, in module order.""" + if self._injector is None: + raise RuntimeError( + "Container not built. Call container.build() before resolving providers." + ) + + instances: List[Any] = [] + seen: set[int] = set() + for module_ref in self._modules.values(): + for descriptor in module_ref.compiled.provider_descriptors: + instance = self.get(descriptor.provide) + instance_id = id(instance) + if instance_id in seen: + continue + seen.add(instance_id) + instances.append(instance) + return instances + + def get_instances_of(self, base_type: Type[T]) -> List[T]: + """Return registered provider instances matching ``base_type``.""" + return [ + instance + for instance in self.get_provider_instances() + if isinstance(instance, base_type) + ] + def clear(self) -> None: """Reset container state. Useful in tests.""" self._injector = None diff --git a/nest/core/pynest_factory.py b/nest/core/pynest_factory.py index a74b24f..464d08b 100644 --- a/nest/core/pynest_factory.py +++ b/nest/core/pynest_factory.py @@ -23,7 +23,12 @@ class PyNestFactory(AbstractPyNestFactory): """Factory that creates a fully-wired PyNest application from a root module.""" @staticmethod - def create(main_module: Type[ModuleType], **kwargs) -> PyNestApp: + def create( + main_module: Type[ModuleType], + *, + worker_grace_timeout: float = 10.0, + **kwargs, + ) -> PyNestApp: """ Build and return a PyNestApp. @@ -39,7 +44,11 @@ def create(main_module: Type[ModuleType], **kwargs) -> PyNestApp: PyNestFactory._run_async(container.initialize_lifecycle()) http_server = FastAPI(**kwargs) - return PyNestApp(container, http_server) + return PyNestApp( + container, + http_server, + worker_grace_timeout=worker_grace_timeout, + ) @staticmethod def _create_server(**kwargs) -> FastAPI: diff --git a/nest/core/worker_factory.py b/nest/core/worker_factory.py new file mode 100644 index 0000000..97c01c6 --- /dev/null +++ b/nest/core/worker_factory.py @@ -0,0 +1,204 @@ +from __future__ import annotations + +import asyncio +import signal as signal_module +from typing import Iterable, Optional, Type + +from nest.core.pynest_container import PyNestContainer +from nest.core.pynest_factory import AbstractPyNestFactory, ModuleType +from nest.core.worker_host import WorkerHost + + +class WorkerApplication: + """A dependency-injection application that runs without an HTTP server.""" + + def __init__( + self, + container: PyNestContainer, + *, + grace_timeout: float = 10.0, + ) -> None: + self.container = container + self._worker_host = WorkerHost.discover(container, grace_timeout=grace_timeout) + self._shutdown_requested = asyncio.Event() + self._started = False + self._closed = False + self._close_task: Optional[asyncio.Task[None]] = None + self.signal: Optional[str] = None + + @property + def closed(self) -> bool: + return self._closed + + def get_worker_host(self) -> WorkerHost: + return self._worker_host + + async def start(self) -> None: + """Initialize application lifecycle and start worker tasks.""" + if self._closed: + raise RuntimeError("Cannot start a closed worker application") + if self._started: + return + + await self.container.initialize_lifecycle() + await self._worker_host.start() + self._started = True + + def request_shutdown(self, signal: Optional[str] = None) -> None: + """Request shutdown without blocking a signal-handler callback.""" + if signal is not None: + self.signal = signal + self._shutdown_requested.set() + + async def run(self) -> None: + """Run until all workers finish or process shutdown is requested.""" + await self.start() + workers_done = asyncio.create_task( + self._worker_host.wait(), name="pynest-workers-complete" + ) + shutdown_requested = asyncio.create_task( + self._shutdown_requested.wait(), name="pynest-workers-shutdown-request" + ) + + try: + await asyncio.wait( + (workers_done, shutdown_requested), + return_when=asyncio.FIRST_COMPLETED, + ) + finally: + try: + await self.close(self.signal) + finally: + for task in (workers_done, shutdown_requested): + if not task.done(): + task.cancel() + await asyncio.gather( + workers_done, shutdown_requested, return_exceptions=True + ) + + async def close(self, signal: Optional[str] = None) -> None: + """Stop workers and run container shutdown hooks once.""" + if self._close_task is not None: + await asyncio.shield(self._close_task) + return + if self._closed: + return + + if signal is not None: + self.signal = signal + self._close_task = asyncio.create_task( + self._close(), name="pynest-worker-application-close" + ) + await asyncio.shield(self._close_task) + + async def _close(self) -> None: + try: + await self._worker_host.stop() + if self._started: + await self.container.shutdown_lifecycle(self.signal) + self._closed = True + finally: + self._started = False + + +class WorkerAppFactory(AbstractPyNestFactory): + """Build standalone, non-HTTP PyNest worker applications.""" + + @staticmethod + def create( + main_module: Type[ModuleType], + *, + grace_timeout: float = 10.0, + **kwargs, + ) -> WorkerApplication: + if kwargs: + unexpected = ", ".join(sorted(kwargs)) + raise TypeError(f"Unexpected worker application options: {unexpected}") + + container = PyNestContainer() + container.add_module(main_module) + container.build() + return WorkerApplication(container, grace_timeout=grace_timeout) + + +def run_workers( + main_module: Type[ModuleType], + *, + grace_timeout: float = 10.0, + signals: Optional[Iterable[signal_module.Signals]] = None, +) -> None: + """Run a standalone worker application until completion or a signal.""" + try: + asyncio.get_running_loop() + except RuntimeError: + pass + else: + raise RuntimeError( + "run_workers() cannot be called from a running event loop; " + "use WorkerAppFactory.create() and await app.run() instead" + ) + + asyncio.run( + _run_workers( + main_module, + grace_timeout=grace_timeout, + signals=signals, + ) + ) + + +async def _run_workers( + main_module: Type[ModuleType], + *, + grace_timeout: float, + signals: Optional[Iterable[signal_module.Signals]], +) -> None: + app = WorkerAppFactory.create(main_module, grace_timeout=grace_timeout) + loop = asyncio.get_running_loop() + shutdown_signals = ( + tuple(signals) + if signals is not None + else (signal_module.SIGTERM, signal_module.SIGINT) + ) + previous_handlers = { + shutdown_signal: signal_module.getsignal(shutdown_signal) + for shutdown_signal in shutdown_signals + } + loop_handlers: set[signal_module.Signals] = set() + + try: + for shutdown_signal in shutdown_signals: + try: + loop.add_signal_handler( + shutdown_signal, + app.request_shutdown, + shutdown_signal.name, + ) + loop_handlers.add(shutdown_signal) + except (NotImplementedError, RuntimeError): + signal_module.signal( + shutdown_signal, + _make_signal_handler(loop, app, shutdown_signal), + ) + + await app.run() + finally: + for shutdown_signal, previous_handler in previous_handlers.items(): + if shutdown_signal in loop_handlers: + loop.remove_signal_handler(shutdown_signal) + signal_module.signal(shutdown_signal, previous_handler) + + +def _make_signal_handler( + loop: asyncio.AbstractEventLoop, + app: WorkerApplication, + shutdown_signal: signal_module.Signals, +): + def handler(signum, frame) -> None: + try: + signal_name = signal_module.Signals(signum).name + except ValueError: + signal_name = shutdown_signal.name + loop.call_soon_threadsafe(app.request_shutdown, signal_name) + + return handler diff --git a/nest/core/worker_host.py b/nest/core/worker_host.py new file mode 100644 index 0000000..91f10ae --- /dev/null +++ b/nest/core/worker_host.py @@ -0,0 +1,340 @@ +from __future__ import annotations + +import asyncio +import inspect +import logging +import math +from dataclasses import dataclass +from typing import Any, Iterable, Optional + +from nest.common.background_worker import ( + BackgroundWorker, + RestartPolicy, + WorkerState, + WorkerStatus, +) +from nest.core.pynest_container import PyNestContainer + + +@dataclass +class _WorkerRuntime: + worker: BackgroundWorker + name: str + restart: RestartPolicy + restart_backoff: float + restart_backoff_max: float + state: WorkerState = WorkerState.STOPPED + restarts: int = 0 + last_error: Optional[str] = None + task: Optional[asyncio.Task[None]] = None + start_hook_called: bool = False + stop_hook_called: bool = False + + +class WorkerHost: + """Supervises all background worker providers for one application.""" + + def __init__( + self, + workers: Iterable[BackgroundWorker], + *, + grace_timeout: float = 10.0, + logger: Optional[logging.Logger] = None, + ) -> None: + self._grace_timeout = self._validate_number(grace_timeout, "grace_timeout") + self._logger = logger or logging.getLogger("pynest.workers") + self._runtimes = self._build_runtimes(workers) + self._running = False + self._stop_task: Optional[asyncio.Task[None]] = None + + @classmethod + def discover( + cls, + container: PyNestContainer, + *, + grace_timeout: float = 10.0, + ) -> "WorkerHost": + """Create a host from dependency-injected worker providers.""" + return cls( + container.get_instances_of(BackgroundWorker), + grace_timeout=grace_timeout, + ) + + @property + def workers(self) -> tuple[BackgroundWorker, ...]: + return tuple(runtime.worker for runtime in self._runtimes) + + @property + def running(self) -> bool: + return self._running + + async def start(self) -> None: + """Start every worker on the current event loop.""" + if self._running: + return + + self._running = True + self._stop_task = None + for runtime in self._runtimes: + runtime.worker._prepare_for_start() + runtime.state = WorkerState.STARTING + runtime.restarts = 0 + runtime.last_error = None + runtime.start_hook_called = False + runtime.stop_hook_called = False + runtime.task = asyncio.create_task( + self._supervise(runtime), + name=f"pynest-worker:{runtime.name}", + ) + + async def stop(self) -> None: + """Request cooperative shutdown, then cancel work past the deadline.""" + if self._stop_task is not None: + await asyncio.shield(self._stop_task) + return + if not self._running: + return + + self._stop_task = asyncio.create_task( + self._stop(), name="pynest-worker-shutdown" + ) + await asyncio.shield(self._stop_task) + + async def wait(self) -> None: + """Wait until all currently managed workers have exited.""" + tasks = [runtime.task for runtime in self._runtimes if runtime.task is not None] + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + + def get_status(self) -> tuple[WorkerStatus, ...]: + """Return immutable status snapshots in discovery order.""" + return tuple( + WorkerStatus( + name=runtime.name, + state=runtime.state, + restarts=runtime.restarts, + last_error=runtime.last_error, + ) + for runtime in self._runtimes + ) + + def status(self) -> list[dict[str, object]]: + """Return JSON-ready status dictionaries.""" + return [status.as_dict() for status in self.get_status()] + + async def _supervise(self, runtime: _WorkerRuntime) -> None: + if runtime.worker.stopping.is_set(): + runtime.state = WorkerState.STOPPED + return + + try: + await self._call_hook(runtime.worker.on_start) + except asyncio.CancelledError: + runtime.state = WorkerState.STOPPED + raise + except Exception as exc: + self._record_failure(runtime, exc, "start") + return + runtime.start_hook_called = True + + if runtime.worker.stopping.is_set(): + await self._run_stop_hook(runtime) + runtime.state = WorkerState.STOPPED + return + + while not runtime.worker.stopping.is_set(): + runtime.state = WorkerState.RUNNING + failed = False + try: + await runtime.worker.run() + except asyncio.CancelledError: + runtime.state = ( + WorkerState.STOPPED + if runtime.worker.stopping.is_set() + else WorkerState.FAILED + ) + raise + except Exception as exc: + failed = True + self._record_failure(runtime, exc, "run") + + if runtime.worker.stopping.is_set(): + runtime.state = WorkerState.STOPPED + return + + should_restart = runtime.restart is RestartPolicy.ALWAYS or ( + failed and runtime.restart is RestartPolicy.ON_FAILURE + ) + if not should_restart: + runtime.state = WorkerState.FAILED if failed else WorkerState.COMPLETED + return + + runtime.restarts += 1 + runtime.state = WorkerState.BACKING_OFF + delay = min( + runtime.restart_backoff * (2 ** min(runtime.restarts - 1, 62)), + runtime.restart_backoff_max, + ) + self._logger.info( + "Restarting background worker %s in %.3f seconds", + runtime.name, + delay, + ) + if not await runtime.worker.sleep(delay): + runtime.state = WorkerState.STOPPED + return + + runtime.state = WorkerState.STOPPED + + async def _stop(self) -> None: + loop = asyncio.get_running_loop() + deadline = loop.time() + self._grace_timeout + active_runtimes = [ + runtime + for runtime in self._runtimes + if runtime.task is not None and not runtime.task.done() + ] + hook_runtimes = [ + runtime + for runtime in self._runtimes + if runtime.start_hook_called and not runtime.stop_hook_called + ] + + try: + for runtime in active_runtimes: + runtime.state = WorkerState.STOPPING + runtime.worker._request_stop() + + hook_tasks = [ + asyncio.create_task( + self._run_stop_hook(runtime), + name=f"pynest-worker-stop-hook:{runtime.name}", + ) + for runtime in hook_runtimes + ] + await self._wait_until_deadline(hook_tasks, deadline) + + worker_tasks = [ + runtime.task for runtime in active_runtimes if runtime.task is not None + ] + await self._wait_until_deadline(worker_tasks, deadline) + finally: + for runtime in active_runtimes: + if runtime.state is WorkerState.STOPPING: + runtime.state = WorkerState.STOPPED + self._running = False + + async def _run_stop_hook(self, runtime: _WorkerRuntime) -> None: + if runtime.stop_hook_called: + return + runtime.stop_hook_called = True + try: + await self._call_hook(runtime.worker.on_stop) + except asyncio.CancelledError: + raise + except Exception as exc: + if runtime.last_error is None: + runtime.last_error = self._format_error(exc) + self._logger.exception( + "Background worker %s failed during stop", runtime.name + ) + + async def _wait_until_deadline( + self, tasks: Iterable[asyncio.Task[Any]], deadline: float + ) -> None: + task_set = set(tasks) + if not task_set: + return + + remaining = max(0.0, deadline - asyncio.get_running_loop().time()) + _, pending = await asyncio.wait(task_set, timeout=remaining) + for task in pending: + task.cancel() + if pending: + await asyncio.gather(*pending, return_exceptions=True) + + def _record_failure( + self, runtime: _WorkerRuntime, exc: Exception, phase: str + ) -> None: + runtime.last_error = self._format_error(exc) + runtime.state = WorkerState.FAILED + self._logger.exception( + "Background worker %s failed during %s", runtime.name, phase + ) + + def _build_runtimes( + self, workers: Iterable[BackgroundWorker] + ) -> list[_WorkerRuntime]: + runtimes: list[_WorkerRuntime] = [] + names: set[str] = set() + + for worker in workers: + if not isinstance(worker, BackgroundWorker): + raise TypeError( + "WorkerHost only accepts BackgroundWorker instances, " + f"got {type(worker).__name__}" + ) + name = self._worker_name(worker) + if name in names: + raise ValueError(f"Duplicate background worker name: {name!r}") + names.add(name) + + try: + restart = RestartPolicy(worker.restart) + except (TypeError, ValueError) as exc: + raise ValueError( + f"Invalid restart policy for background worker {name!r}: " + f"{worker.restart!r}" + ) from exc + + restart_backoff = self._validate_number( + worker.restart_backoff, + f"restart_backoff for background worker {name!r}", + ) + restart_backoff_max = self._validate_number( + worker.restart_backoff_max, + f"restart_backoff_max for background worker {name!r}", + ) + if restart_backoff_max < restart_backoff: + raise ValueError( + f"restart_backoff_max for background worker {name!r} " + "must be greater than or equal to restart_backoff" + ) + + runtimes.append( + _WorkerRuntime( + worker=worker, + name=name, + restart=restart, + restart_backoff=restart_backoff, + restart_backoff_max=restart_backoff_max, + ) + ) + + return runtimes + + @staticmethod + async def _call_hook(hook) -> None: + result = hook() + if inspect.isawaitable(result): + await result + + @staticmethod + def _worker_name(worker: BackgroundWorker) -> str: + configured_name = worker.name + if not isinstance(configured_name, str): + raise ValueError("Background worker name must be a string") + return configured_name.strip() or type(worker).__name__ + + @staticmethod + def _validate_number(value: float, name: str) -> float: + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ValueError(f"{name} must be a finite, non-negative number") + number = float(value) + if not math.isfinite(number) or number < 0: + raise ValueError(f"{name} must be a finite, non-negative number") + return number + + @staticmethod + def _format_error(exc: BaseException) -> str: + return f"{type(exc).__name__}: {exc}" diff --git a/tests/test_common/test_background_worker.py b/tests/test_common/test_background_worker.py new file mode 100644 index 0000000..9e1f381 --- /dev/null +++ b/tests/test_common/test_background_worker.py @@ -0,0 +1,167 @@ +import asyncio + +import pytest + +from nest.common.background_worker import ( + BackgroundWorker, + IntervalWorker, + RestartPolicy, + WorkerState, + WorkerStatus, +) + + +class ProbeWorker(BackgroundWorker): + async def run(self) -> None: + pass + + +def test_background_worker_has_safe_restart_defaults(): + worker = ProbeWorker() + + assert worker.name == "" + assert worker.restart is RestartPolicy.ON_FAILURE + assert worker.restart_backoff == 1.0 + assert worker.restart_backoff_max == 30.0 + assert worker.stopping.is_set() is False + + +def test_worker_does_not_require_subclasses_to_call_super_init(): + class WorkerWithDependencies(BackgroundWorker): + def __init__(self, dependency): + self.dependency = dependency + + async def run(self) -> None: + pass + + worker = WorkerWithDependencies(object()) + + assert worker.stopping.is_set() is False + + +def test_sleep_returns_true_after_the_delay_elapses(): + async def scenario(): + worker = ProbeWorker() + + assert await worker.sleep(0) is True + + asyncio.run(scenario()) + + +def test_sleep_returns_false_when_stop_is_requested(): + async def scenario(): + worker = ProbeWorker() + task = asyncio.create_task(worker.sleep(10)) + await asyncio.sleep(0) + + worker.stopping.set() + + assert await asyncio.wait_for(task, timeout=0.1) is False + + asyncio.run(scenario()) + + +def test_sleep_rejects_negative_delays(): + async def scenario(): + worker = ProbeWorker() + + with pytest.raises(ValueError, match="greater than or equal to zero"): + await worker.sleep(-0.1) + + asyncio.run(scenario()) + + +def test_prepare_for_start_replaces_a_previous_stop_event(): + worker = ProbeWorker() + old_event = worker.stopping + old_event.set() + + worker._prepare_for_start() + + assert worker.stopping is not old_event + assert worker.stopping.is_set() is False + + +def test_interval_worker_runs_immediately_and_never_overlaps_executions(): + class ProbeIntervalWorker(IntervalWorker): + interval = 0.001 + run_immediately = True + + def __init__(self): + self.calls = 0 + self.in_flight = 0 + self.max_in_flight = 0 + + async def execute(self) -> None: + self.calls += 1 + self.in_flight += 1 + self.max_in_flight = max(self.max_in_flight, self.in_flight) + await asyncio.sleep(0) + self.in_flight -= 1 + if self.calls == 2: + self.stopping.set() + + async def scenario(): + worker = ProbeIntervalWorker() + await worker.run() + + assert worker.calls == 2 + assert worker.max_in_flight == 1 + + asyncio.run(scenario()) + + +def test_interval_worker_waits_before_first_execution_by_default(): + class DelayedWorker(IntervalWorker): + interval = 10 + + def __init__(self): + self.calls = 0 + + async def execute(self) -> None: + self.calls += 1 + + async def scenario(): + worker = DelayedWorker() + task = asyncio.create_task(worker.run()) + await asyncio.sleep(0) + + assert worker.calls == 0 + + worker.stopping.set() + await asyncio.wait_for(task, timeout=0.1) + assert worker.calls == 0 + + asyncio.run(scenario()) + + +def test_interval_worker_rejects_a_negative_interval(): + class InvalidIntervalWorker(IntervalWorker): + interval = -1 + + async def execute(self) -> None: + pass + + async def scenario(): + with pytest.raises(ValueError, match="greater than zero"): + await InvalidIntervalWorker().run() + + asyncio.run(scenario()) + + +def test_worker_status_is_immutable_and_json_ready(): + status = WorkerStatus( + name="email", + state=WorkerState.BACKING_OFF, + restarts=2, + last_error="connection lost", + ) + + assert status.as_dict() == { + "name": "email", + "state": "backing_off", + "restarts": 2, + "last_error": "connection lost", + } + with pytest.raises(AttributeError): + status.restarts = 3 diff --git a/tests/test_core/test_worker_host.py b/tests/test_core/test_worker_host.py new file mode 100644 index 0000000..20734db --- /dev/null +++ b/tests/test_core/test_worker_host.py @@ -0,0 +1,325 @@ +import asyncio + +import pytest + +from nest.common.background_worker import ( + BackgroundWorker, + RestartPolicy, + WorkerState, +) +from nest.core import Injectable, Module +from nest.core.pynest_container import PyNestContainer +from nest.core.worker_host import WorkerHost + + +def test_container_discovers_each_worker_provider_once(): + @Injectable + class TestWorker(BackgroundWorker): + async def run(self) -> None: + pass + + @Module(providers=[TestWorker]) + class WorkerModule: + pass + + @Module(imports=[WorkerModule], providers=[TestWorker]) + class AppModule: + pass + + container = PyNestContainer() + container.add_module(AppModule) + container.build() + + assert container.get_instances_of(BackgroundWorker) == [container.get(TestWorker)] + + +def test_get_instances_of_requires_a_built_container(): + container = PyNestContainer() + + with pytest.raises(RuntimeError, match="build"): + container.get_instances_of(BackgroundWorker) + + +def test_host_starts_and_cooperatively_stops_a_worker(): + class CooperativeWorker(BackgroundWorker): + name = "cooperative" + + def __init__(self): + self.started = asyncio.Event() + self.on_start_calls = 0 + self.on_stop_calls = 0 + self.exited = False + + async def on_start(self) -> None: + self.on_start_calls += 1 + + async def run(self) -> None: + self.started.set() + await self.stopping.wait() + self.exited = True + + async def on_stop(self) -> None: + self.on_stop_calls += 1 + + async def scenario(): + worker = CooperativeWorker() + host = WorkerHost([worker], grace_timeout=0.1) + + await host.start() + await asyncio.wait_for(worker.started.wait(), timeout=0.1) + assert host.status() == [ + { + "name": "cooperative", + "state": "running", + "restarts": 0, + "last_error": None, + } + ] + + await host.stop() + await host.stop() + + assert worker.on_start_calls == 1 + assert worker.on_stop_calls == 1 + assert worker.exited is True + assert host.get_status()[0].state is WorkerState.STOPPED + + asyncio.run(scenario()) + + +def test_host_restarts_a_failed_worker_with_backoff(): + class FailsOnceWorker(BackgroundWorker): + restart_backoff = 0 + + def __init__(self): + self.attempts = 0 + self.succeeded = asyncio.Event() + + async def run(self) -> None: + self.attempts += 1 + if self.attempts == 1: + raise RuntimeError("temporary failure") + self.succeeded.set() + await self.stopping.wait() + + async def scenario(): + worker = FailsOnceWorker() + host = WorkerHost([worker], grace_timeout=0.1) + + await host.start() + await asyncio.wait_for(worker.succeeded.wait(), timeout=0.1) + + status = host.get_status()[0] + assert worker.attempts == 2 + assert status.restarts == 1 + assert status.last_error == "RuntimeError: temporary failure" + assert status.state is WorkerState.RUNNING + + await host.stop() + + asyncio.run(scenario()) + + +def test_none_policy_leaves_a_crashed_worker_failed(): + class FailingWorker(BackgroundWorker): + restart = RestartPolicy.NONE + + async def run(self) -> None: + raise LookupError("permanent failure") + + async def scenario(): + host = WorkerHost([FailingWorker()]) + + await host.start() + await asyncio.wait_for(host.wait(), timeout=0.1) + + assert host.status() == [ + { + "name": "FailingWorker", + "state": "failed", + "restarts": 0, + "last_error": "LookupError: permanent failure", + } + ] + + await host.stop() + + asyncio.run(scenario()) + + +def test_always_policy_restarts_a_normally_completed_worker(): + class RepeatWorker(BackgroundWorker): + restart = RestartPolicy.ALWAYS + restart_backoff = 0 + + def __init__(self): + self.calls = 0 + self.restarted = asyncio.Event() + + async def run(self) -> None: + self.calls += 1 + if self.calls == 2: + self.restarted.set() + await self.stopping.wait() + + async def scenario(): + worker = RepeatWorker() + host = WorkerHost([worker], grace_timeout=0.1) + + await host.start() + await asyncio.wait_for(worker.restarted.wait(), timeout=0.1) + + assert host.get_status()[0].restarts == 1 + await host.stop() + + asyncio.run(scenario()) + + +def test_host_cancels_a_worker_after_the_grace_timeout(): + class StubbornWorker(BackgroundWorker): + def __init__(self): + self.started = asyncio.Event() + self.never = asyncio.Event() + self.cancelled = False + + async def run(self) -> None: + self.started.set() + try: + await self.never.wait() + except asyncio.CancelledError: + self.cancelled = True + raise + + async def scenario(): + worker = StubbornWorker() + host = WorkerHost([worker], grace_timeout=0.01) + + await host.start() + await asyncio.wait_for(worker.started.wait(), timeout=0.1) + await host.stop() + + assert worker.cancelled is True + assert host.get_status()[0].state is WorkerState.STOPPED + + asyncio.run(scenario()) + + +def test_stop_timeout_also_bounds_slow_on_stop_hooks(): + class SlowStopWorker(BackgroundWorker): + def __init__(self): + self.started = asyncio.Event() + self.hook_cancelled = False + + async def run(self) -> None: + self.started.set() + await self.stopping.wait() + + async def on_stop(self) -> None: + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + self.hook_cancelled = True + raise + + async def scenario(): + worker = SlowStopWorker() + host = WorkerHost([worker], grace_timeout=0.01) + + await host.start() + await asyncio.wait_for(worker.started.wait(), timeout=0.1) + await host.stop() + + assert worker.hook_cancelled is True + + asyncio.run(scenario()) + + +def test_stop_waits_for_start_hook_before_calling_stop_hook(): + class SlowStartWorker(BackgroundWorker): + def __init__(self): + self.starting = asyncio.Event() + self.release_start = asyncio.Event() + self.events = [] + + async def on_start(self) -> None: + self.events.append("start:begin") + self.starting.set() + await self.release_start.wait() + self.events.append("start:end") + + async def run(self) -> None: + self.events.append("run") + + async def on_stop(self) -> None: + self.events.append("stop") + + async def scenario(): + worker = SlowStartWorker() + host = WorkerHost([worker], grace_timeout=0.1) + await host.start() + await asyncio.wait_for(worker.starting.wait(), timeout=0.1) + + stop_task = asyncio.create_task(host.stop()) + await asyncio.sleep(0) + worker.release_start.set() + await asyncio.wait_for(stop_task, timeout=0.1) + + assert worker.events == ["start:begin", "start:end", "stop"] + + asyncio.run(scenario()) + + +def test_duplicate_worker_names_are_rejected(): + class FirstWorker(BackgroundWorker): + name = "duplicate" + + async def run(self) -> None: + pass + + class SecondWorker(BackgroundWorker): + name = "duplicate" + + async def run(self) -> None: + pass + + with pytest.raises(ValueError, match="Duplicate background worker name"): + WorkerHost([FirstWorker(), SecondWorker()]) + + +def test_blank_worker_name_falls_back_to_the_class_name(): + class BlankNameWorker(BackgroundWorker): + name = " " + + async def run(self) -> None: + pass + + assert WorkerHost([BlankNameWorker()]).status()[0]["name"] == "BlankNameWorker" + + +def test_host_rejects_values_that_are_not_workers(): + with pytest.raises(TypeError, match="BackgroundWorker"): + WorkerHost([object()]) + + +@pytest.mark.parametrize( + ("attribute", "value", "message"), + [ + ("restart", "sometimes", "restart policy"), + ("restart_backoff", -1, "restart_backoff"), + ("restart_backoff_max", -1, "restart_backoff_max"), + ], +) +def test_invalid_worker_configuration_is_rejected(attribute, value, message): + class InvalidWorker(BackgroundWorker): + async def run(self) -> None: + pass + + setattr(InvalidWorker, attribute, value) + + with pytest.raises(ValueError, match=message): + WorkerHost([InvalidWorker()]) + + +def test_negative_grace_timeout_is_rejected(): + with pytest.raises(ValueError, match="grace_timeout"): + WorkerHost([], grace_timeout=-1) diff --git a/tests/test_core/test_worker_integration.py b/tests/test_core/test_worker_integration.py new file mode 100644 index 0000000..de0b4a8 --- /dev/null +++ b/tests/test_core/test_worker_integration.py @@ -0,0 +1,238 @@ +import asyncio +import threading + +import pytest +from fastapi.testclient import TestClient + +from nest.common import OnApplicationBootstrap, OnModuleDestroy +from nest.common.background_worker import BackgroundWorker +from nest.core import Injectable, Module, PyNestFactory +from nest.core.worker_factory import WorkerAppFactory, run_workers + + +def test_http_lifespan_runs_workers_and_stops_them_before_provider_shutdown(): + events = [] + + @Injectable + class Dependency(OnModuleDestroy): + def on_module_destroy(self) -> None: + events.append("dependency:destroy") + + @Injectable + class HttpWorker(BackgroundWorker): + def __init__(self, dependency: Dependency): + self.dependency = dependency + self.started = threading.Event() + + async def run(self) -> None: + events.append("worker:start") + self.started.set() + await self.stopping.wait() + events.append("worker:exit") + + async def on_stop(self) -> None: + events.append("worker:stop") + + @Module(providers=[HttpWorker, Dependency]) + class AppModule: + pass + + app = PyNestFactory.create(AppModule) + worker = app.container.get(HttpWorker) + + assert worker.dependency is app.container.get(Dependency) + assert worker.started.is_set() is False + + with TestClient(app.get_server()): + assert worker.started.wait(timeout=1) + assert app.get_worker_host().status() == [ + { + "name": "HttpWorker", + "state": "running", + "restarts": 0, + "last_error": None, + } + ] + + assert events.index("worker:exit") < events.index("dependency:destroy") + assert app.get_worker_host().status()[0]["state"] == "stopped" + + +def test_concurrent_http_close_callers_wait_for_the_same_shutdown(): + destroy_started = asyncio.Event() + allow_destroy = asyncio.Event() + + @Injectable + class SlowDependency(OnModuleDestroy): + async def on_module_destroy(self) -> None: + destroy_started.set() + await allow_destroy.wait() + + @Module(providers=[SlowDependency]) + class AppModule: + pass + + app = PyNestFactory.create(AppModule) + + async def scenario(): + first_close = asyncio.create_task(app.close()) + await asyncio.wait_for(destroy_started.wait(), timeout=0.1) + second_close = asyncio.create_task(app.close()) + await asyncio.sleep(0) + + assert second_close.done() is False + + allow_destroy.set() + await asyncio.gather(first_close, second_close) + + asyncio.run(scenario()) + + +def test_standalone_app_runs_lifecycle_and_worker_on_one_loop(): + events = [] + loop_ids = [] + + @Injectable + class BootstrapProbe(OnApplicationBootstrap, OnModuleDestroy): + async def on_application_bootstrap(self) -> None: + events.append("bootstrap") + loop_ids.append(id(asyncio.get_running_loop())) + + async def on_module_destroy(self) -> None: + events.append("destroy") + loop_ids.append(id(asyncio.get_running_loop())) + + @Injectable + class CompletingWorker(BackgroundWorker): + async def run(self) -> None: + events.append("worker") + loop_ids.append(id(asyncio.get_running_loop())) + + @Module(providers=[BootstrapProbe, CompletingWorker]) + class AppModule: + pass + + async def scenario(): + app = WorkerAppFactory.create(AppModule, grace_timeout=0.1) + + await app.run() + + assert events == ["bootstrap", "worker", "destroy"] + assert len(set(loop_ids)) == 1 + assert app.closed is True + + asyncio.run(scenario()) + + +def test_standalone_request_shutdown_stops_a_long_running_worker(): + @Injectable + class LongRunningWorker(BackgroundWorker): + def __init__(self): + self.started = asyncio.Event() + self.exited = False + + async def run(self) -> None: + self.started.set() + await self.stopping.wait() + self.exited = True + + @Module(providers=[LongRunningWorker]) + class AppModule: + pass + + async def scenario(): + app = WorkerAppFactory.create(AppModule, grace_timeout=0.1) + worker = app.container.get(LongRunningWorker) + run_task = asyncio.create_task(app.run()) + await asyncio.wait_for(worker.started.wait(), timeout=0.1) + + app.request_shutdown("SIGTERM") + await asyncio.wait_for(run_task, timeout=0.2) + + assert worker.exited is True + assert app.signal == "SIGTERM" + assert app.get_worker_host().status()[0]["state"] == "stopped" + + asyncio.run(scenario()) + + +def test_standalone_close_is_idempotent(): + @Injectable + class CompletingWorker(BackgroundWorker): + async def run(self) -> None: + pass + + @Module(providers=[CompletingWorker]) + class AppModule: + pass + + async def scenario(): + app = WorkerAppFactory.create(AppModule) + + await app.run() + await app.close() + + assert app.closed is True + + asyncio.run(scenario()) + + +def test_cancelling_standalone_run_still_closes_the_application(): + @Injectable + class LongRunningWorker(BackgroundWorker): + def __init__(self): + self.started = asyncio.Event() + self.exited = False + + async def run(self) -> None: + self.started.set() + await self.stopping.wait() + self.exited = True + + @Module(providers=[LongRunningWorker]) + class AppModule: + pass + + async def scenario(): + app = WorkerAppFactory.create(AppModule, grace_timeout=0.1) + worker = app.container.get(LongRunningWorker) + run_task = asyncio.create_task(app.run()) + await asyncio.wait_for(worker.started.wait(), timeout=0.1) + + run_task.cancel() + with pytest.raises(asyncio.CancelledError): + await run_task + + assert app.closed is True + assert worker.exited is True + + asyncio.run(scenario()) + + +def test_run_workers_runs_a_completing_worker(): + calls = [] + + @Injectable + class CompletingWorker(BackgroundWorker): + async def run(self) -> None: + calls.append("ran") + + @Module(providers=[CompletingWorker]) + class AppModule: + pass + + run_workers(AppModule, signals=[]) + + assert calls == ["ran"] + + +def test_run_workers_rejects_use_inside_a_running_event_loop(): + @Module() + class AppModule: + pass + + async def scenario(): + with pytest.raises(RuntimeError, match="running event loop"): + run_workers(AppModule) + + asyncio.run(scenario())