From bede6f8a4093a32cdb16b5a47d2db6953f8e04a9 Mon Sep 17 00:00:00 2001 From: d3vyce <44915747+d3vyce@users.noreply.github.com> Date: Mon, 4 May 2026 18:53:11 +0200 Subject: [PATCH] fix: skip deduplication for non-JSON-serializable kwargs instead of raising TypeError (#12) --- src/taskiq_deduplication/middleware.py | 34 +++++++++++++++++--------- tests/test_middleware.py | 15 +++++++++++- 2 files changed, 37 insertions(+), 12 deletions(-) diff --git a/src/taskiq_deduplication/middleware.py b/src/taskiq_deduplication/middleware.py index 1056cb3..8559f28 100644 --- a/src/taskiq_deduplication/middleware.py +++ b/src/taskiq_deduplication/middleware.py @@ -56,7 +56,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): if self._redis is not None: await self._redis.aclose() - def _build_deduplication_key(self, message: TaskiqMessage) -> str: + 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}" @@ -67,10 +67,13 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): if key_fields is not None else message.kwargs ) - payload = json.dumps( - {"task": message.task_name, "kwargs": kwargs}, - sort_keys=True, - ) + try: + payload = json.dumps( + {"task": message.task_name, "kwargs": kwargs}, + sort_keys=True, + ) + except TypeError: + return None fingerprint = hashlib.sha256(payload.encode()).hexdigest()[:16] return f"{self.key_prefix}:{fingerprint}" @@ -94,6 +97,13 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): assert self._redis is not None key = self._build_deduplication_key(message) + if key is None: + logger.warning( + "Task %s has non-JSON-serializable kwargs; deduplication skipped." + " Use the deduplication_key label to deduplicate this task.", + message.task_name, + ) + return message ttl = self._get_ttl(message.labels) logger.debug("Acquiring lock %s for task %s", key, message.task_name) @@ -118,9 +128,10 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): ) -> None: if not self._is_enabled(message.labels): return - await self._release_if_owned( - self._build_deduplication_key(message), message.task_id - ) + key = self._build_deduplication_key(message) + if key is None: + return + await self._release_if_owned(key, message.task_id) async def on_error( self, @@ -130,6 +141,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): ) -> None: if not self._is_enabled(message.labels): return - await self._release_if_owned( - self._build_deduplication_key(message), message.task_id - ) + key = self._build_deduplication_key(message) + if key is None: + return + await self._release_if_owned(key, message.task_id) diff --git a/tests/test_middleware.py b/tests/test_middleware.py index 4254b4e..e454701 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -92,7 +92,7 @@ class TestDefaultBuildDeduplicationKey: mw._redis = None m = make_message() key = mw._build_deduplication_key(m) - assert key.startswith("myapp:locks:") + assert key is not None and key.startswith("myapp:locks:") def test_empty_kwargs_produces_consistent_key(self, middleware, make_message): m1 = make_message(kwargs={}) @@ -128,6 +128,10 @@ class TestDefaultBuildDeduplicationKey: ) assert middleware._build_deduplication_key(m) == "taskiq:deduplication:my-lock" + def test_non_serializable_kwargs_returns_none(self, middleware, make_message): + m = make_message(kwargs={"dt": object()}) + assert middleware._build_deduplication_key(m) is None + class TestPreSend: @pytest.mark.anyio @@ -174,6 +178,15 @@ class TestPreSend: await middleware.pre_send(make_message(kwargs={"x": 1})) await middleware.pre_send(make_message(kwargs={"x": 2})) + @pytest.mark.anyio + async def test_non_serializable_kwargs_skips_deduplication( + self, middleware, make_message + ): + msg1 = make_message(kwargs={"dt": object()}) + msg2 = make_message(kwargs={"dt": object()}) + await middleware.pre_send(msg1) + await middleware.pre_send(msg2) # should not raise + @pytest.mark.anyio async def test_default_ttl_applied(self, fake_redis, make_message): mw = RedisDeduplicationMiddleware(redis_url="redis://localhost", default_ttl=77)