fix: label values stringified by taskiq kicker not parsed correctly (#52)

This commit is contained in:
d3vyce
2026-06-09 18:15:35 +02:00
committed by GitHub
parent 6c5c23fe05
commit 7310950406
3 changed files with 136 additions and 42 deletions
+10 -42
View File
@@ -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,
+55
View File
@@ -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
+71
View File
@@ -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
):