From 7310950406486c4e20021a881dfad553d1197b92 Mon Sep 17 00:00:00 2001 From: d3vyce <44915747+d3vyce@users.noreply.github.com> Date: Tue, 9 Jun 2026 18:15:35 +0200 Subject: [PATCH] fix: label values stringified by taskiq kicker not parsed correctly (#52) --- src/taskiq_deduplication/middleware.py | 52 ++++--------------- src/taskiq_deduplication/utils.py | 55 ++++++++++++++++++++ tests/test_middleware.py | 71 ++++++++++++++++++++++++++ 3 files changed, 136 insertions(+), 42 deletions(-) diff --git a/src/taskiq_deduplication/middleware.py b/src/taskiq_deduplication/middleware.py index 369a1a9..7b90cd0 100644 --- a/src/taskiq_deduplication/middleware.py +++ b/src/taskiq_deduplication/middleware.py @@ -8,7 +8,13 @@ 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 +from .utils import ( + RELEASE_LUA_SCRIPT, + check_and_delete, + parse_bool_label, + parse_int_label, + parse_list_label, +) logger = logging.getLogger(__name__) @@ -91,50 +97,12 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): if self._redis is not None: await self._redis.aclose() - @staticmethod - 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, 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 - - @staticmethod - def _parse_int_label(value: Any, default: int, label_name: str = "") -> int: - try: - return int(value) - except (TypeError, ValueError): - logger.warning( - "Invalid %r value %r (expected int); falling back to default (%d).", - label_name, - value, - default, - ) - return default - 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( + key_fields = parse_list_label( message.labels.get(DEDUP_KEY_FIELDS_LABEL), DEDUP_KEY_FIELDS_LABEL ) kwargs = ( @@ -153,12 +121,12 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): return f"{self.key_prefix}:{fingerprint}" def _is_enabled(self, labels: dict[str, Any]) -> bool: - return self._parse_bool_label( + return parse_bool_label( labels.get(DEDUP_LABEL), self.default_deduplication, DEDUP_LABEL ) def _get_ttl(self, labels: dict[str, Any]) -> int: - return self._parse_int_label( + return parse_int_label( labels.get(DEDUP_TTL_LABEL, self.default_ttl), self.default_ttl, DEDUP_TTL_LABEL, diff --git a/src/taskiq_deduplication/utils.py b/src/taskiq_deduplication/utils.py index dfe7b0a..a366712 100644 --- a/src/taskiq_deduplication/utils.py +++ b/src/taskiq_deduplication/utils.py @@ -1,5 +1,9 @@ +import ast +import logging from typing import Any +logger = logging.getLogger(__name__) + RELEASE_LUA_SCRIPT = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) @@ -22,3 +26,54 @@ async def check_and_delete(script: Any, key: str, owner: str) -> bool: """ released: int = await script(keys=[key], args=[owner]) return bool(released) + + +def parse_bool_label(value: Any, default: bool, label_name: str = "") -> bool: + if isinstance(value, bool): + return value + if isinstance(value, str): + lower = value.lower() + if lower == "true": + return True + if lower == "false": + return False + if value is not None: + logger.warning( + "Invalid %r value %r (expected bool); falling back to default (%r).", + label_name, + value, + default, + ) + return default + + +def parse_list_label(value: Any, label_name: str = "") -> list[str] | None: + if isinstance(value, list): + return value + if isinstance(value, str): + try: + parsed = ast.literal_eval(value) + if isinstance(parsed, list): + return parsed + except (ValueError, SyntaxError): + pass + if value is not None: + logger.warning( + "Invalid %r value %r (expected list[str]); ignoring.", + label_name, + value, + ) + return None + + +def parse_int_label(value: Any, default: int, label_name: str = "") -> int: + try: + return int(value) + except (TypeError, ValueError): + logger.warning( + "Invalid %r value %r (expected int); falling back to default (%d).", + label_name, + value, + default, + ) + return default diff --git a/tests/test_middleware.py b/tests/test_middleware.py index bb3e64a..6fc0fb1 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -480,6 +480,22 @@ class TestLabelTypeParsing: with pytest.raises(DuplicateTaskError): await middleware.pre_send(make_message(labels={DEDUP_LABEL: True})) + async def test_string_bool_label_true_enables_dedup(self, middleware, make_message): + # taskiq's prepare_label() stringifies True → "True" before pre_send runs + 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"})) + + async def test_string_bool_label_false_disables_dedup( + self, fake_redis, make_message + ): + # taskiq's prepare_label() stringifies False → "False" before pre_send runs + 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"})) + async def test_key_fields_list_parsed_correctly(self, middleware, make_message): m1 = make_message( kwargs={"a": 1, "b": 2, "c": 3}, @@ -511,6 +527,20 @@ class TestLabelTypeParsing: await middleware.pre_send(msg) assert any("yes" in r.message for r in caplog.records) + async def test_string_true_lowercase_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"})) + + async def test_string_false_lowercase_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"})) + async def test_invalid_key_fields_label_warns_and_falls_back_to_all_kwargs( self, middleware, make_message, caplog ): @@ -523,6 +553,47 @@ class TestLabelTypeParsing: # falls back to full-kwargs fingerprint — key must still be produced assert key is not None + async def test_key_fields_string_parses_to_non_list_warns( + self, middleware, make_message, caplog + ): + import logging + + # ast.literal_eval succeeds but returns a dict, not a list + msg = make_message(kwargs={"a": 1}, labels={DEDUP_KEY_FIELDS_LABEL: "{'a': 1}"}) + with caplog.at_level(logging.WARNING, logger="taskiq_deduplication.middleware"): + key = middleware._build_deduplication_key(msg) + assert any("{'a': 1}" in r.message for r in caplog.records) + assert key is not None + + def test_stringified_key_fields_parsed_correctly(self, middleware, make_message): + # taskiq's prepare_label() stringifies ["a", "b"] → "['a', 'b']" before pre_send + 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) + + def test_stringified_key_fields_different_included_fields( + self, middleware, make_message + ): + m1 = make_message( + kwargs={"a": 1, "b": 2}, + labels={DEDUP_KEY_FIELDS_LABEL: "['a']"}, + ) + m2 = make_message( + kwargs={"a": 99, "b": 2}, + labels={DEDUP_KEY_FIELDS_LABEL: "['a']"}, + ) + assert middleware._build_deduplication_key( + m1 + ) != middleware._build_deduplication_key(m2) + async def test_invalid_ttl_string_warns_and_uses_default( self, middleware, fake_redis, make_message, caplog ):