mirror of
https://github.com/Viren070/MediaFusion.git
synced 2025-12-01 23:21:11 +01:00
Refactor Dramatiq TaskManager Middleware
Introduced `TaskInfo` dataclass and caching mechanism for task metadata to improve performance. Refactored Redis operations into a dedicated helper method with error handling. Enhanced code readability and reduced redundancy by modularizing task key generation and processing logic.
This commit is contained in:
+98
-54
@@ -3,9 +3,11 @@ import logging
|
||||
import os
|
||||
import signal
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta, datetime
|
||||
from functools import lru_cache
|
||||
from threading import Lock
|
||||
from typing import Callable, Optional
|
||||
from typing import Callable, Optional, Dict
|
||||
|
||||
import dramatiq
|
||||
from apscheduler.triggers.cron import CronTrigger
|
||||
@@ -203,12 +205,24 @@ class Retries(OriginalRetries):
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TaskInfo:
|
||||
task_name: str
|
||||
min_interval: timedelta
|
||||
set_cache_expiry: bool
|
||||
task_key: str
|
||||
|
||||
|
||||
class TaskManager(dramatiq.Middleware):
|
||||
def __init__(self, processing_time_buffer: int = 10):
|
||||
self.processing_time_buffer = processing_time_buffer
|
||||
self._task_info_cache: Dict[str, TaskInfo] = {}
|
||||
|
||||
@staticmethod
|
||||
@lru_cache(maxsize=128)
|
||||
def calculate_interval_from_crontab(crontab_expression: str) -> timedelta:
|
||||
"""
|
||||
Calculate the minimum interval between two consecutive runs
|
||||
Calculate and cache the minimum interval between two consecutive runs
|
||||
specified by a crontab expression.
|
||||
"""
|
||||
cron_trigger = CronTrigger.from_crontab(crontab_expression)
|
||||
@@ -218,7 +232,24 @@ class TaskManager(dramatiq.Middleware):
|
||||
second_next_time = cron_trigger.get_next_fire_time(next_time, next_time)
|
||||
return second_next_time - next_time
|
||||
|
||||
def get_task_data(self, broker, message):
|
||||
def _generate_task_key(self, task_name: str, args: tuple, kwargs: dict) -> str:
|
||||
"""Generate a consistent task key."""
|
||||
if spider_name := kwargs.get("spider_name"):
|
||||
return f"background_tasks:run_spider:spider_name={spider_name}"
|
||||
|
||||
args_str = "_".join(str(arg) for arg in args)
|
||||
kwargs_str = "_".join(f"{k}={v}" for k, v in sorted(kwargs.items()))
|
||||
return f"background_tasks:{task_name}:{args_str}_{kwargs_str}".rstrip("_")
|
||||
|
||||
def get_task_info(self, broker, message) -> Optional[TaskInfo]:
|
||||
"""
|
||||
Get task information with caching for better performance.
|
||||
"""
|
||||
# Try to get from cache first
|
||||
message_id = message.message_id
|
||||
if message_id in self._task_info_cache:
|
||||
return self._task_info_cache[message_id]
|
||||
|
||||
task_name = message.actor_name
|
||||
args = message.args
|
||||
kwargs = message.kwargs.copy()
|
||||
@@ -226,74 +257,87 @@ class TaskManager(dramatiq.Middleware):
|
||||
min_interval = getattr(actor, "_minimum_run_interval", None)
|
||||
set_cache_expiry = False
|
||||
|
||||
if kwargs.get("crontab_expression"):
|
||||
min_interval = self.calculate_interval_from_crontab(
|
||||
kwargs.get("crontab_expression")
|
||||
)
|
||||
if crontab_expr := kwargs.get("crontab_expression"):
|
||||
min_interval = self.calculate_interval_from_crontab(crontab_expr)
|
||||
del kwargs["crontab_expression"]
|
||||
elif min_interval:
|
||||
set_cache_expiry = True
|
||||
else:
|
||||
logging.info(
|
||||
logging.debug(
|
||||
f"No restriction set for task {task_name} with args {args} and kwargs {kwargs}"
|
||||
)
|
||||
return
|
||||
return None
|
||||
|
||||
if spider_name := kwargs.get("spider_name"):
|
||||
task_key = f"background_tasks:run_spider:spider_name={spider_name}"
|
||||
elif video_id := kwargs.get("video_id"):
|
||||
task_key = f"background_tasks:{task_name}:video_id={video_id}"
|
||||
else:
|
||||
keys = "_".join([str(arg) for arg in args])
|
||||
keys += "_".join([f"{k}={v}" for k, v in kwargs.items()])
|
||||
task_key = f"background_tasks:{task_name}:{keys}"
|
||||
task_key = self._generate_task_key(task_name, args, kwargs)
|
||||
task_info = TaskInfo(task_name, min_interval, set_cache_expiry, task_key)
|
||||
|
||||
return task_name, min_interval, set_cache_expiry, task_key
|
||||
# Cache the result
|
||||
self._task_info_cache[message_id] = task_info
|
||||
return task_info
|
||||
|
||||
def _check_and_update_redis(
|
||||
self, task_info: TaskInfo, operation: str = "check"
|
||||
) -> Optional[bool]:
|
||||
"""
|
||||
Handle Redis operations with error handling and logging.
|
||||
"""
|
||||
try:
|
||||
if operation == "check":
|
||||
last_run = REDIS_SYNC_CLIENT.get(task_info.task_key)
|
||||
if last_run is not None:
|
||||
last_run = datetime.fromtimestamp(float(last_run))
|
||||
difference = datetime.now() - last_run
|
||||
min_interval = task_info.min_interval - timedelta(
|
||||
seconds=self.processing_time_buffer
|
||||
)
|
||||
|
||||
if difference < min_interval:
|
||||
logging.warning(
|
||||
f"Discarding task {task_info.task_name} with task_key {task_info.task_key}. "
|
||||
f"Last run: {difference} ago. Minimum interval: {min_interval}"
|
||||
)
|
||||
return True
|
||||
|
||||
# Update Redis with new timestamp
|
||||
ex_time = (
|
||||
int(task_info.min_interval.total_seconds())
|
||||
if task_info.set_cache_expiry
|
||||
else None
|
||||
)
|
||||
REDIS_SYNC_CLIENT.set(
|
||||
task_info.task_key,
|
||||
datetime.now().timestamp(),
|
||||
ex=ex_time,
|
||||
)
|
||||
logging.debug(f"Task key {task_info.task_key} updated in Redis")
|
||||
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"Redis operation failed for task {task_info.task_key}: {str(e)}"
|
||||
)
|
||||
# Don't skip the message if Redis fails
|
||||
return False
|
||||
|
||||
return None
|
||||
|
||||
def before_process_message(self, broker, message):
|
||||
task_data = self.get_task_data(broker, message)
|
||||
if not task_data:
|
||||
return
|
||||
task_name, min_interval, set_cache_expiry, task_key = task_data
|
||||
|
||||
# Subtract 10 seconds to account for processing time
|
||||
min_interval = min_interval - timedelta(seconds=10)
|
||||
|
||||
last_run = REDIS_SYNC_CLIENT.get(task_key)
|
||||
if last_run is not None:
|
||||
last_run = datetime.fromtimestamp(float(last_run))
|
||||
difference = datetime.now() - last_run
|
||||
if difference < min_interval:
|
||||
logging.warning(
|
||||
f"Discarding task {task_name} with task_key {task_key} due to minimum run interval. Last run: {difference} ago. Minimum interval: {min_interval}"
|
||||
)
|
||||
if task_info := self.get_task_info(broker, message):
|
||||
if self._check_and_update_redis(task_info, "check"):
|
||||
raise SkipMessage()
|
||||
|
||||
# Set the cache expiry for the task
|
||||
ex_time = int(min_interval.total_seconds()) if set_cache_expiry else None
|
||||
REDIS_SYNC_CLIENT.set(
|
||||
task_key,
|
||||
datetime.now().timestamp(),
|
||||
ex=ex_time,
|
||||
)
|
||||
logging.info(f"Executing task {task_name} with task key {task_key}")
|
||||
|
||||
def after_process_message(self, broker, message, *, result=None, exception=None):
|
||||
if exception:
|
||||
return
|
||||
|
||||
task_data = self.get_task_data(broker, message)
|
||||
if not task_data:
|
||||
return
|
||||
task_name, min_interval, set_cache_expiry, task_key = task_data
|
||||
if task_info := self.get_task_info(broker, message):
|
||||
self._check_and_update_redis(task_info, "update")
|
||||
|
||||
# Update the cache with the latest run time
|
||||
REDIS_SYNC_CLIENT.set(
|
||||
task_key,
|
||||
datetime.now().timestamp(),
|
||||
ex=int(min_interval.total_seconds()) if set_cache_expiry else None,
|
||||
)
|
||||
logging.info(f"Task key {task_key} updated cache with latest run time.")
|
||||
# Cleanup cache
|
||||
self._task_info_cache.pop(message.message_id, None)
|
||||
|
||||
def after_skip_message(self, broker, message):
|
||||
# Cleanup cache for skipped messages
|
||||
self._task_info_cache.pop(message.message_id, None)
|
||||
|
||||
|
||||
class TimingMiddleware(BaseHTTPMiddleware):
|
||||
|
||||
Reference in New Issue
Block a user