diff --git a/src/taskiq_deduplication/middleware.py b/src/taskiq_deduplication/middleware.py index 4ed5c74..3886f20 100644 --- a/src/taskiq_deduplication/middleware.py +++ b/src/taskiq_deduplication/middleware.py @@ -90,15 +90,28 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): await self._redis.aclose() @staticmethod - def _parse_bool_label(value: Any, default: bool) -> bool: + def _parse_bool_label(value: Any, default: bool, label_name: str = "") -> bool: if isinstance(value, bool): return value + if value is not None: + logger.warning( + "Invalid %r value %r (expected bool); falling back to default (%r).", + label_name, + value, + default, + ) return default @staticmethod - def _parse_list_label(value: Any) -> list[str] | None: + def _parse_list_label(value: Any, label_name: str = "") -> list[str] | None: if isinstance(value, list): return value + if value is not None: + logger.warning( + "Invalid %r value %r (expected list[str]); ignoring.", + label_name, + value, + ) return None def _build_deduplication_key(self, message: TaskiqMessage) -> str | None: @@ -106,7 +119,9 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): 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)) + key_fields = self._parse_list_label( + message.labels.get(DEDUP_KEY_FIELDS_LABEL), 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 @@ -124,11 +139,20 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): def _is_enabled(self, labels: dict[str, Any]) -> bool: return self._parse_bool_label( - labels.get(DEDUP_LABEL), self.default_deduplication + labels.get(DEDUP_LABEL), self.default_deduplication, DEDUP_LABEL ) def _get_ttl(self, labels: dict[str, Any]) -> int: - return int(labels.get(DEDUP_TTL_LABEL, self.default_ttl)) + value = labels.get(DEDUP_TTL_LABEL, self.default_ttl) + try: + return int(value) + except (TypeError, ValueError): + logger.warning( + "Invalid deduplication_ttl value %r; falling back to default (%ds).", + value, + self.default_ttl, + ) + return self.default_ttl async def _release_if_owned(self, key: str, task_id: str) -> None: if self._redis is None: diff --git a/tests/test_middleware.py b/tests/test_middleware.py index 3087397..9b33cc1 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -494,6 +494,58 @@ class TestLabelTypeParsing: key = middleware._build_deduplication_key(m) assert key is not None + @pytest.mark.anyio + async def test_invalid_bool_label_warns_and_uses_default( + self, middleware, make_message, caplog + ): + import logging + + msg = make_message(labels={DEDUP_LABEL: "yes"}) + with caplog.at_level(logging.WARNING, logger="taskiq_deduplication.middleware"): + await middleware.pre_send(msg) + assert any("yes" in r.message for r in caplog.records) + + @pytest.mark.anyio + async def test_invalid_key_fields_label_warns_and_falls_back_to_all_kwargs( + self, middleware, make_message, caplog + ): + import logging + + msg = make_message(kwargs={"a": 1}, labels={DEDUP_KEY_FIELDS_LABEL: "user_id"}) + with caplog.at_level(logging.WARNING, logger="taskiq_deduplication.middleware"): + key = middleware._build_deduplication_key(msg) + assert any("user_id" in r.message for r in caplog.records) + # falls back to full-kwargs fingerprint — key must still be produced + assert key is not None + + @pytest.mark.anyio + async def test_invalid_ttl_string_warns_and_uses_default( + self, middleware, fake_redis, make_message, caplog + ): + import logging + + msg = make_message(labels={DEDUP_TTL_LABEL: "oops"}) + with caplog.at_level(logging.WARNING, logger="taskiq_deduplication.middleware"): + await middleware.pre_send(msg) + assert any("oops" in r.message for r in caplog.records) + key = middleware._build_deduplication_key(msg) + ttl = await fake_redis.ttl(key) + assert 0 < ttl <= middleware.default_ttl + + @pytest.mark.anyio + async def test_invalid_ttl_none_warns_and_uses_default( + self, middleware, fake_redis, make_message, caplog + ): + import logging + + msg = make_message(labels={DEDUP_TTL_LABEL: None}) + with caplog.at_level(logging.WARNING, logger="taskiq_deduplication.middleware"): + await middleware.pre_send(msg) + assert any("None" in r.message for r in caplog.records) + key = middleware._build_deduplication_key(msg) + ttl = await fake_redis.ttl(key) + assert 0 < ttl <= middleware.default_ttl + class TestKeyCaching: @pytest.mark.anyio