Merge pull request #81 from arabcoders/dev

wrap notifier in timeout call, and expose pool operations via API
This commit is contained in:
Abdulmohsen 2024-03-23 02:14:31 +03:00 committed by GitHub
commit fbaa013cff
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 118 additions and 27 deletions

View file

@ -39,20 +39,19 @@ class AsyncPool:
self._num_workers = num_workers self._num_workers = num_workers
self._logger = logger if logger else logging.getLogger(__name__) self._logger = logger if logger else logging.getLogger(__name__)
self._queue = asyncio.Queue(num_workers * load_factor) self._queue = asyncio.Queue(num_workers * load_factor)
self._workers = None self._workers: dict[str, asyncio.Future] = {}
self._exceptions = False self._exceptions = False
self._max_task_time = max_task_time self._max_task_time = max_task_time
self._return_futures = return_futures self._return_futures = return_futures
self._raise_on_join = raise_on_join self._raise_on_join = raise_on_join
self._name = name self._name = name
self._worker_co = worker_co self._worker_co = worker_co
self._active_worker: dict[str, dict | None] = {} self._status: dict[str, dict | None] = {}
async def _worker_loop(self, worker_id: str): async def _worker_loop(self, worker_id: str):
while True: while True:
got_obj = False got_obj = False
future = None future = None
try: try:
item = await self._queue.get() item = await self._queue.get()
got_obj = True got_obj = True
@ -62,7 +61,7 @@ class AsyncPool:
future, args, kwargs = item future, args, kwargs = item
self._active_worker[worker_id] = { self._status[worker_id] = {
'started': self._time().isoformat(), 'started': self._time().isoformat(),
'data': kwargs.get('entry', {'info': {}}).info.__dict__, 'data': kwargs.get('entry', {'info': {}}).info.__dict__,
} }
@ -79,6 +78,9 @@ class AsyncPool:
self._exceptions = True self._exceptions = True
raise raise
except BaseException as e: except BaseException as e:
if isinstance(e, asyncio.exceptions.CancelledError):
raise e
self._exceptions = True self._exceptions = True
if future: if future:
@ -87,7 +89,7 @@ class AsyncPool:
else: else:
self._logger.exception(f'Worker call failed. {str(e)}') self._logger.exception(f'Worker call failed. {str(e)}')
finally: finally:
self._active_worker[worker_id] = None self._status[worker_id] = None
if got_obj: if got_obj:
self._queue.task_done() self._queue.task_done()
@ -106,13 +108,13 @@ class AsyncPool:
""" """
:return: number of available workers. :return: number of available workers.
""" """
return sum(1 for worker_status in self._active_worker.values() if worker_status is None) return sum(1 for worker_status in self._status.values() if worker_status is None)
def get_workers_status(self) -> dict[str, datetime | None]: def get_workers_status(self) -> dict[str, datetime | None]:
""" """
:return: dictionary of worker status. :return: dictionary of worker status.
""" """
return self._active_worker return self._status
async def __aenter__(self): async def __aenter__(self):
self.start() self.start()
@ -139,23 +141,55 @@ class AsyncPool:
def start(self): def start(self):
""" Will start up worker pool and reset exception state """ """ Will start up worker pool and reset exception state """
assert self._workers is None
self._exceptions = False self._exceptions = False
self._workers = []
for worker_number in range(self._num_workers): for worker_number in range(self._num_workers):
workerName = f"worker_{worker_number+1}" worker_id = f"worker_{worker_number+1}"
self._workers.append( self._createWorker(worker_id)
asyncio.ensure_future(
coro_or_future=self._worker_loop(worker_id=workerName),
loop=self._loop
)
)
self._active_worker[workerName] = None
async def join(self): async def restart(self, worker_id: str, msg: str = None) -> bool:
""" Will restart the worker pool """
if worker_id not in self._workers:
self._logger.warning(f'Worker {worker_id} does not exist.')
return False
try:
self._workers[worker_id].cancel(msg)
await self._workers[worker_id]
except asyncio.exceptions.CancelledError as e:
self._logger.warning(f'Worker {worker_id} restarted. {str(e)}')
if worker_id in self._status:
self._status.pop(worker_id)
if worker_id in self._workers:
self._workers.pop(worker_id)
self._createWorker(worker_id)
return True
async def stop(self, worker_id: str, msg: str = None) -> bool:
""" Will stop the worker """
if worker_id not in self._workers:
self._logger.warning(f'Worker {worker_id} does not exist.')
return False
if self._workers[worker_id].cancel(msg):
try:
await self._workers[worker_id]
except asyncio.exceptions.CancelledError as e:
self._logger.warning(f'Worker {worker_id} stopped. {str(e)}')
if worker_id in self._status:
self._status.pop(worker_id)
if worker_id in self._workers:
self._workers.pop(worker_id)
return True
async def join(self) -> None:
# no-op if workers aren't running # no-op if workers aren't running
if not self._workers: if len(self._workers) < 1:
return return
self._logger.info(f'Joining {self._name}') self._logger.info(f'Joining {self._name}')
@ -166,8 +200,8 @@ class AsyncPool:
await self._queue.put(Terminator()) await self._queue.put(Terminator())
try: try:
await asyncio.gather(*self._workers) await asyncio.gather(*list(self._workers))
self._workers = None self._workers = {}
except: except:
self._logger.exception(f'Exception joining {self._name}') self._logger.exception(f'Exception joining {self._name}')
raise raise
@ -177,5 +211,20 @@ class AsyncPool:
if self._exceptions and self._raise_on_join: if self._exceptions and self._raise_on_join:
raise Exception(f"Exception occurred in {self._name} pool") raise Exception(f"Exception occurred in {self._name} pool")
def _time(self): def _createWorker(self, worker_id: str) -> asyncio.Future:
if worker_id in self._workers:
self._logger.debug(f'Worker {worker_id} already exists.')
return self._workers[worker_id]
self._status[worker_id] = None
self._workers[worker_id] = asyncio.ensure_future(
coro_or_future=self._worker_loop(worker_id=worker_id),
loop=self._loop
)
self._logger.debug(f'Created {worker_id}')
return self._workers[worker_id]
def _time(self) -> datetime:
return datetime.now(tz=timezone.utc) return datetime.now(tz=timezone.utc)

View file

@ -11,10 +11,10 @@ from Webhooks import Webhooks
LOG = logging.getLogger('Utils') LOG = logging.getLogger('Utils')
AUDIO_FORMATS: tuple = ('m4a', 'mp3', 'opus', 'wav') AUDIO_FORMATS: tuple[str] = ('m4a', 'mp3', 'opus', 'wav',)
IGNORED_KEYS: tuple = ('cookiefile', 'paths', 'outtmpl', 'progress_hooks', 'postprocessor_hooks',) IGNORED_KEYS: tuple[str] = ('cookiefile', 'paths', 'outtmpl', 'progress_hooks', 'postprocessor_hooks',)
YTDLP_INFO_CLS = None YTDLP_INFO_CLS: yt_dlp.YoutubeDL = None
def get_format(format: str, quality: str) -> str: def get_format(format: str, quality: str) -> str:
@ -347,4 +347,7 @@ class Notifier:
if self.webhooks and event in self.webhooks_allowed_events: if self.webhooks and event in self.webhooks_allowed_events:
tasks.append(self.webhooks.send(event, data)) tasks.append(self.webhooks.send(event, data))
await asyncio.gather(*tasks) try:
await asyncio.wait_for(asyncio.gather(*tasks), timeout=60)
except asyncio.TimeoutError:
LOG.error(f"Timed out sending event {event} to webhooks.")

View file

@ -364,10 +364,49 @@ class Main:
def default(o): return f"<<non-serializable: {type(o).__qualname__}>>" def default(o): return f"<<non-serializable: {type(o).__qualname__}>>"
return json.dumps(obj, default=default) return json.dumps(obj, default=default)
return web.Response(text=safe_serialize(data), headers={ return web.Response(text=safe_serialize({
'open': self.queue.pool.has_open_workers(),
'count': self.queue.pool.get_available_workers(),
'workers': data,
}), headers={
'Content-Type': 'application/json', 'Content-Type': 'application/json',
}) })
@self.routes.post(self.config.url_prefix + 'workers')
async def restart_pool(_) -> Response:
if self.queue.pool is None:
return web.json_response({"error": "Worker pool not initialized."}, status=404)
status = self.queue.pool.start()
return web.Response()
@self.routes.patch(self.config.url_prefix + 'workers/{id}')
async def restart_worker(request: Request) -> Response:
id: str = request.match_info.get('id')
if not id:
raise web.HTTPBadRequest(reason='worker id is required.')
if self.queue.pool is None:
return web.json_response({"error": "Worker pool not initialized."}, status=404)
status = await self.queue.pool.restart(id, 'requested by user.')
return web.json_response({"status": "restarted" if status else "in_error_state"})
@self.routes.delete(self.config.url_prefix + 'workers/{id}')
async def stop_worker(request: Request) -> Response:
id: str = request.match_info.get('id')
if not id:
raise web.HTTPBadRequest(reason='worker id is required.')
if self.queue.pool is None:
return web.json_response({"error": "Worker pool not initialized."}, status=404)
status = await self.queue.pool.stop(id, 'requested by user.')
return web.json_response({"status": "stopped" if status else "in_error_state"})
@self.routes.get(self.config.url_prefix + 'm3u8/{file:.*}') @self.routes.get(self.config.url_prefix + 'm3u8/{file:.*}')
async def m3u8(request: Request) -> Response: async def m3u8(request: Request) -> Response:
file: str = request.match_info.get('file') file: str = request.match_info.get('file')