mirror of
https://github.com/d3vyce/taskiq-deduplication.git
synced 2026-08-04 19:14:07 +00:00
fix: label values stringified by taskiq kicker not parsed correctly (#52)
This commit is contained in:
@@ -8,7 +8,13 @@ from redis.asyncio import Redis
|
|||||||
from taskiq import TaskiqMessage, TaskiqResult
|
from taskiq import TaskiqMessage, TaskiqResult
|
||||||
from taskiq.abc.middleware import TaskiqMiddleware
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -91,50 +97,12 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
if self._redis is not None:
|
if self._redis is not None:
|
||||||
await self._redis.aclose()
|
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:
|
def _build_deduplication_key(self, message: TaskiqMessage) -> str | None:
|
||||||
explicit_key: str | None = message.labels.get(DEDUP_EXPLICIT_KEY_LABEL)
|
explicit_key: str | None = message.labels.get(DEDUP_EXPLICIT_KEY_LABEL)
|
||||||
if explicit_key is not None:
|
if explicit_key is not None:
|
||||||
return f"{self.key_prefix}:{explicit_key}"
|
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
|
message.labels.get(DEDUP_KEY_FIELDS_LABEL), DEDUP_KEY_FIELDS_LABEL
|
||||||
)
|
)
|
||||||
kwargs = (
|
kwargs = (
|
||||||
@@ -153,12 +121,12 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
return f"{self.key_prefix}:{fingerprint}"
|
return f"{self.key_prefix}:{fingerprint}"
|
||||||
|
|
||||||
def _is_enabled(self, labels: dict[str, Any]) -> bool:
|
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
|
labels.get(DEDUP_LABEL), self.default_deduplication, DEDUP_LABEL
|
||||||
)
|
)
|
||||||
|
|
||||||
def _get_ttl(self, labels: dict[str, Any]) -> int:
|
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),
|
labels.get(DEDUP_TTL_LABEL, self.default_ttl),
|
||||||
self.default_ttl,
|
self.default_ttl,
|
||||||
DEDUP_TTL_LABEL,
|
DEDUP_TTL_LABEL,
|
||||||
|
|||||||
@@ -1,5 +1,9 @@
|
|||||||
|
import ast
|
||||||
|
import logging
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
RELEASE_LUA_SCRIPT = """
|
RELEASE_LUA_SCRIPT = """
|
||||||
if redis.call('get', KEYS[1]) == ARGV[1] then
|
if redis.call('get', KEYS[1]) == ARGV[1] then
|
||||||
return redis.call('del', KEYS[1])
|
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])
|
released: int = await script(keys=[key], args=[owner])
|
||||||
return bool(released)
|
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
|
||||||
|
|||||||
@@ -480,6 +480,22 @@ class TestLabelTypeParsing:
|
|||||||
with pytest.raises(DuplicateTaskError):
|
with pytest.raises(DuplicateTaskError):
|
||||||
await middleware.pre_send(make_message(labels={DEDUP_LABEL: True}))
|
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):
|
async def test_key_fields_list_parsed_correctly(self, middleware, make_message):
|
||||||
m1 = make_message(
|
m1 = make_message(
|
||||||
kwargs={"a": 1, "b": 2, "c": 3},
|
kwargs={"a": 1, "b": 2, "c": 3},
|
||||||
@@ -511,6 +527,20 @@ class TestLabelTypeParsing:
|
|||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
assert any("yes" in r.message for r in caplog.records)
|
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(
|
async def test_invalid_key_fields_label_warns_and_falls_back_to_all_kwargs(
|
||||||
self, middleware, make_message, caplog
|
self, middleware, make_message, caplog
|
||||||
):
|
):
|
||||||
@@ -523,6 +553,47 @@ class TestLabelTypeParsing:
|
|||||||
# falls back to full-kwargs fingerprint — key must still be produced
|
# falls back to full-kwargs fingerprint — key must still be produced
|
||||||
assert key is not None
|
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(
|
async def test_invalid_ttl_string_warns_and_uses_default(
|
||||||
self, middleware, fake_redis, make_message, caplog
|
self, middleware, fake_redis, make_message, caplog
|
||||||
):
|
):
|
||||||
|
|||||||
Reference in New Issue
Block a user