|
| 1 | +import asyncio |
1 | 2 | import os |
2 | 3 | import sys |
3 | | -from collections.abc import Generator, Sequence |
| 4 | +from collections.abc import Callable, Generator, Sequence |
4 | 5 | from contextlib import contextmanager |
5 | 6 | from importlib import import_module |
6 | 7 | from logging import getLogger |
|
9 | 10 |
|
10 | 11 | logger = getLogger("taskiq.worker") |
11 | 12 |
|
| 13 | +LoopFactory = Callable[[], asyncio.AbstractEventLoop] |
| 14 | + |
12 | 15 |
|
13 | 16 | @contextmanager |
14 | 17 | def add_cwd_in_path() -> Generator[None, None, None]: |
@@ -55,6 +58,38 @@ def import_object(object_spec: str, app_dir: str | None = None) -> Any: |
55 | 58 | return getattr(module, import_spec[1]) |
56 | 59 |
|
57 | 60 |
|
| 61 | +def resolve_loop_factory( |
| 62 | + loop_factory: str, |
| 63 | + app_dir: str | None = None, |
| 64 | +) -> LoopFactory: |
| 65 | + """ |
| 66 | + Resolve an event loop factory from a callable or import string. |
| 67 | +
|
| 68 | + :param loop_factory: path in `module:variable` format. |
| 69 | + :param app_dir: directory to add in sys.path for importing. |
| 70 | + :raises ValueError: if the resolved object is not callable. |
| 71 | + :return: event loop factory. |
| 72 | + """ |
| 73 | + factory = import_object(loop_factory, app_dir=app_dir) |
| 74 | + if not callable(factory): |
| 75 | + raise ValueError("Event loop factory must be callable.") |
| 76 | + return factory |
| 77 | + |
| 78 | + |
| 79 | +def create_event_loop(loop_factory: LoopFactory) -> asyncio.AbstractEventLoop: |
| 80 | + """ |
| 81 | + Create and validate an event loop from a factory. |
| 82 | +
|
| 83 | + :param loop_factory: event loop factory. |
| 84 | + :raises ValueError: if the factory does not return an event loop. |
| 85 | + :return: created event loop. |
| 86 | + """ |
| 87 | + loop = loop_factory() |
| 88 | + if not isinstance(loop, asyncio.AbstractEventLoop): |
| 89 | + raise ValueError("Event loop factory must return an event loop.") |
| 90 | + return loop |
| 91 | + |
| 92 | + |
58 | 93 | def import_from_modules(modules: list[str]) -> None: |
59 | 94 | """ |
60 | 95 | Import all modules from modules variable. |
|
0 commit comments