mirror of
https://github.com/d3vyce/taskiq-deduplication.git
synced 2026-08-04 19:14:07 +00:00
refactor: simplify label handling and add key caching (#22)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user