import asyncio import hashlib import json import logging from typing import Any, Awaitable, cast from redis.asyncio import Redis from taskiq import TaskiqMessage, TaskiqResult from taskiq.abc.middleware import TaskiqMiddleware from .utils import RELEASE_LUA_SCRIPT, check_and_delete logger = logging.getLogger(__name__) DEDUP_LABEL = "deduplication" DEDUP_TTL_LABEL = "deduplication_ttl" DEDUP_KEY_FIELDS_LABEL = "deduplication_key_fields" DEDUP_EXPLICIT_KEY_LABEL = "deduplication_key" _CACHED_KEY_LABEL = "__taskiq_dedup_cached_key" class DuplicateTaskError(Exception): """Raised when a task with identical name and kwargs is already queued or running.""" class RedisDeduplicationMiddleware(TaskiqMiddleware): """Prevents duplicate tasks from being queued. When a task is dispatched, a Redis lock is acquired for the duration of its execution. Any subsequent task with the same fingerprint is rejected with ``DuplicateTaskError`` while the lock is held. The lock is released automatically on completion or error. Attributes: redis_url: Redis connection URL passed to ``Redis.from_url``. default_deduplication: Whether deduplication is enabled by default. default_ttl: Default lock TTL in seconds. key_prefix: Prefix for all Redis lock keys. """ def __init__( self, redis_url: str, default_deduplication: bool = True, default_ttl: int = 300, key_prefix: str = "taskiq:deduplication", startup_retries: int = 3, startup_retry_delay: float = 1.0, ) -> None: self.redis_url = redis_url self.default_deduplication = default_deduplication self.default_ttl = default_ttl self.key_prefix = key_prefix self.startup_retries = startup_retries self.startup_retry_delay = startup_retry_delay self._redis: Redis | None = None self._release_script: Any = None async def startup(self) -> None: last_error: BaseException | None = None for attempt in range(self.startup_retries): try: self._redis = Redis.from_url(self.redis_url) await cast(Awaitable[bool], self._redis.ping()) self._release_script = self._redis.register_script(RELEASE_LUA_SCRIPT) return except Exception as exc: last_error = exc if attempt < self.startup_retries - 1: delay = self.startup_retry_delay * (2**attempt) logger.warning( "Failed to connect to Redis (attempt %d/%d): %s. " "Retrying in %.1fs...", attempt + 1, self.startup_retries, exc, delay, ) await asyncio.sleep(delay) logger.error( "Failed to connect to Redis after %d attempts.", self.startup_retries ) raise ConnectionError( f"Could not connect to Redis after {self.startup_retries} attempts" ) from last_error async def shutdown(self) -> None: if self._redis is not None: await self._redis.aclose() @staticmethod def _parse_bool_label(value: Any, default: bool) -> bool: if isinstance(value, bool): return value return default @staticmethod def _parse_list_label(value: Any) -> list[str] | None: if isinstance(value, list): return value return None def _build_deduplication_key(self, message: TaskiqMessage) -> str | None: explicit_key: str | None = message.labels.get(DEDUP_EXPLICIT_KEY_LABEL) if explicit_key is not None: return f"{self.key_prefix}:{explicit_key}" key_fields = self._parse_list_label(message.labels.get(DEDUP_KEY_FIELDS_LABEL)) kwargs = ( {k: v for k, v in message.kwargs.items() if k in key_fields} if key_fields is not None else message.kwargs ) try: payload = json.dumps( {"task": message.task_name, "kwargs": kwargs}, sort_keys=True, ) except TypeError: return None fingerprint = hashlib.sha256(payload.encode()).hexdigest()[:16] return f"{self.key_prefix}:{fingerprint}" def _is_enabled(self, labels: dict[str, Any]) -> bool: return self._parse_bool_label( labels.get(DEDUP_LABEL), self.default_deduplication ) def _get_ttl(self, labels: dict[str, Any]) -> int: return int(labels.get(DEDUP_TTL_LABEL, self.default_ttl)) async def _release_if_owned(self, key: str, task_id: str) -> None: if self._redis is None: raise RuntimeError( "RedisDeduplicationMiddleware.startup() was never called." ) if self._release_script is None: self._release_script = self._redis.register_script(RELEASE_LUA_SCRIPT) released = await check_and_delete(self._release_script, key, task_id) if released: logger.debug("Released lock %s", key) else: logger.debug("Skipped release of lock %s: not owned by this task", key) @staticmethod def _get_cached_key(message: TaskiqMessage) -> str | None: return message.labels.get(_CACHED_KEY_LABEL) @staticmethod def _cache_key(message: TaskiqMessage, key: str | None) -> None: message.labels[_CACHED_KEY_LABEL] = key async def pre_send(self, message: TaskiqMessage) -> TaskiqMessage: if not self._is_enabled(message.labels): return message if self._redis is None: raise RuntimeError( "RedisDeduplicationMiddleware.startup() was never called." ) key = self._build_deduplication_key(message) self._cache_key(message, key) if key is None: logger.warning( "Task %s has non-JSON-serializable kwargs; deduplication skipped." " Use the deduplication_key label to deduplicate this task.", message.task_name, ) return message ttl = self._get_ttl(message.labels) logger.debug("Acquiring lock %s for task %s", key, message.task_name) acquired = await self._redis.set(key, message.task_id, ex=ttl, nx=True) if not acquired: logger.warning( "Duplicate task %s dropped (key=%s).", message.task_name, key, ) raise DuplicateTaskError( f"Task {message.task_name!r} with the same arguments is already queued or running." ) logger.debug("Lock %s acquired for task %s", key, message.task_name) return message async def _release_lock(self, message: TaskiqMessage) -> None: if not self._is_enabled(message.labels): return key = self._get_cached_key(message) if key is None: return await self._release_if_owned(key, message.task_id) async def post_execute( self, message: TaskiqMessage, result: TaskiqResult, ) -> None: await self._release_lock(message) async def on_error( self, message: TaskiqMessage, result: TaskiqResult, exception: BaseException, ) -> None: await self._release_lock(message)