diff --git a/src/taskiq_deduplication/middleware.py b/src/taskiq_deduplication/middleware.py index ff5b663..e29a4dc 100644 --- a/src/taskiq_deduplication/middleware.py +++ b/src/taskiq_deduplication/middleware.py @@ -17,6 +17,8 @@ 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.""" @@ -85,12 +87,24 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): 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: list[str] | None = message.labels.get(DEDUP_KEY_FIELDS_LABEL) + 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 @@ -107,7 +121,9 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): return f"{self.key_prefix}:{fingerprint}" def _is_enabled(self, labels: dict[str, Any]) -> bool: - return bool(labels.get(DEDUP_LABEL, self.default_deduplication)) + 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)) @@ -123,6 +139,14 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): 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 @@ -132,6 +156,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): "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." @@ -163,7 +188,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): ) -> None: if not self._is_enabled(message.labels): return - key = self._build_deduplication_key(message) + key = self._get_cached_key(message) if key is None: return await self._release_if_owned(key, message.task_id) @@ -176,7 +201,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): ) -> None: if not self._is_enabled(message.labels): return - key = self._build_deduplication_key(message) + key = self._get_cached_key(message) if key is None: return await self._release_if_owned(key, message.task_id) diff --git a/tests/test_middleware.py b/tests/test_middleware.py index 2dba6e3..22f402f 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -404,3 +404,108 @@ class TestStartupRetry: assert mock_sleep.call_count == 2 assert mock_sleep.call_args_list[0].args[0] == 0.01 assert mock_sleep.call_args_list[1].args[0] == 0.02 + + +class TestLabelTypeParsing: + @pytest.mark.anyio + async def test_bool_label_false_disables_dedup(self, fake_redis, make_message): + mw = RedisDeduplicationMiddleware(redis_url="redis://localhost") + mw._redis = fake_redis + await mw.pre_send(make_message(labels={DEDUP_LABEL: False})) + await mw.pre_send(make_message(labels={DEDUP_LABEL: False})) + + @pytest.mark.anyio + async def test_bool_label_true_enables_dedup(self, middleware, make_message): + msg = make_message(labels={DEDUP_LABEL: True}) + await middleware.pre_send(msg) + with pytest.raises(DuplicateTaskError): + await middleware.pre_send(make_message(labels={DEDUP_LABEL: True})) + + @pytest.mark.anyio + async def test_key_fields_list_parsed_correctly(self, middleware, make_message): + m1 = make_message( + kwargs={"a": 1, "b": 2, "c": 3}, + labels={DEDUP_KEY_FIELDS_LABEL: ["a", "b"]}, + ) + m2 = make_message( + kwargs={"a": 1, "b": 2, "c": 999}, + labels={DEDUP_KEY_FIELDS_LABEL: ["a", "b"]}, + ) + assert middleware._build_deduplication_key( + m1 + ) == middleware._build_deduplication_key(m2) + + @pytest.mark.anyio + async def test_key_fields_non_list_ignored(self, middleware, make_message): + m = make_message( + kwargs={"a": 1}, + labels={DEDUP_KEY_FIELDS_LABEL: "not-a-list"}, + ) + key = middleware._build_deduplication_key(m) + assert key is not None + + +class TestKeyCaching: + @pytest.mark.anyio + async def test_key_cached_during_pre_send(self, middleware, make_message): + from taskiq_deduplication.middleware import _CACHED_KEY_LABEL + + msg = make_message() + await middleware.pre_send(msg) + assert _CACHED_KEY_LABEL in msg.labels + + @pytest.mark.anyio + async def test_post_execute_uses_cached_key( + self, middleware, fake_redis, make_message, make_result + ): + from taskiq_deduplication.middleware import _CACHED_KEY_LABEL + + msg = make_message() + await middleware.pre_send(msg) + cached_key = msg.labels.get(_CACHED_KEY_LABEL) + await middleware.post_execute(msg, make_result()) + assert not await fake_redis.exists(cached_key) + + @pytest.mark.anyio + async def test_on_error_uses_cached_key( + self, middleware, fake_redis, make_message, make_result + ): + from taskiq_deduplication.middleware import _CACHED_KEY_LABEL + + msg = make_message() + await middleware.pre_send(msg) + cached_key = msg.labels.get(_CACHED_KEY_LABEL) + await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom")) + assert not await fake_redis.exists(cached_key) + + @pytest.mark.anyio + async def test_cached_key_none_when_key_build_fails(self, middleware, make_message): + from taskiq_deduplication.middleware import _CACHED_KEY_LABEL + + msg = make_message(kwargs={"dt": object()}) + await middleware.pre_send(msg) + assert msg.labels[_CACHED_KEY_LABEL] is None + + @pytest.mark.anyio + async def test_post_execute_noop_when_cached_key_is_none( + self, middleware, fake_redis, make_message, make_result + ): + msg = make_message(kwargs={"dt": object()}) + await middleware.pre_send(msg) + await middleware.post_execute(msg, make_result()) + + @pytest.mark.anyio + async def test_on_error_noop_when_cached_key_is_none( + self, middleware, fake_redis, make_message, make_result + ): + msg = make_message(kwargs={"dt": object()}) + await middleware.pre_send(msg) + await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom")) + + @pytest.mark.anyio + async def test_release_if_owned_raises_without_redis( + self, middleware, make_message + ): + middleware._redis = None + with pytest.raises(RuntimeError, match="startup"): + await middleware._release_if_owned("some-key", "some-task")