fix: skip deduplication for non-JSON-serializable kwargs instead of raising TypeError (#12)

This commit is contained in:
d3vyce
2026-05-04 18:53:11 +02:00
committed by GitHub
parent 446071924e
commit bede6f8a40
2 changed files with 37 additions and 12 deletions
+23 -11
View File
@@ -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)
+14 -1
View File
@@ -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)