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:
mhdzumair
2024-12-15 22:58:40 +05:30
parent f2712a6cfb
commit d163109ce1
+98 -54
View File
@@ -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):