mirror of
https://github.com/d3vyce/taskiq-deduplication.git
synced 2026-08-04 19:14:07 +00:00
fix: skip deduplication for non-JSON-serializable kwargs instead of raising TypeError (#12)
This commit is contained in:
@@ -56,7 +56,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
if self._redis is not None:
|
if self._redis is not None:
|
||||||
await self._redis.aclose()
|
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)
|
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}"
|
||||||
@@ -67,10 +67,13 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
if key_fields is not None
|
if key_fields is not None
|
||||||
else message.kwargs
|
else message.kwargs
|
||||||
)
|
)
|
||||||
payload = json.dumps(
|
try:
|
||||||
{"task": message.task_name, "kwargs": kwargs},
|
payload = json.dumps(
|
||||||
sort_keys=True,
|
{"task": message.task_name, "kwargs": kwargs},
|
||||||
)
|
sort_keys=True,
|
||||||
|
)
|
||||||
|
except TypeError:
|
||||||
|
return None
|
||||||
fingerprint = hashlib.sha256(payload.encode()).hexdigest()[:16]
|
fingerprint = hashlib.sha256(payload.encode()).hexdigest()[:16]
|
||||||
return f"{self.key_prefix}:{fingerprint}"
|
return f"{self.key_prefix}:{fingerprint}"
|
||||||
|
|
||||||
@@ -94,6 +97,13 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
|
|
||||||
assert self._redis is not None
|
assert self._redis is not None
|
||||||
key = self._build_deduplication_key(message)
|
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)
|
ttl = self._get_ttl(message.labels)
|
||||||
|
|
||||||
logger.debug("Acquiring lock %s for task %s", key, message.task_name)
|
logger.debug("Acquiring lock %s for task %s", key, message.task_name)
|
||||||
@@ -118,9 +128,10 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
) -> None:
|
) -> None:
|
||||||
if not self._is_enabled(message.labels):
|
if not self._is_enabled(message.labels):
|
||||||
return
|
return
|
||||||
await self._release_if_owned(
|
key = self._build_deduplication_key(message)
|
||||||
self._build_deduplication_key(message), message.task_id
|
if key is None:
|
||||||
)
|
return
|
||||||
|
await self._release_if_owned(key, message.task_id)
|
||||||
|
|
||||||
async def on_error(
|
async def on_error(
|
||||||
self,
|
self,
|
||||||
@@ -130,6 +141,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
) -> None:
|
) -> None:
|
||||||
if not self._is_enabled(message.labels):
|
if not self._is_enabled(message.labels):
|
||||||
return
|
return
|
||||||
await self._release_if_owned(
|
key = self._build_deduplication_key(message)
|
||||||
self._build_deduplication_key(message), message.task_id
|
if key is None:
|
||||||
)
|
return
|
||||||
|
await self._release_if_owned(key, message.task_id)
|
||||||
|
|||||||
@@ -92,7 +92,7 @@ class TestDefaultBuildDeduplicationKey:
|
|||||||
mw._redis = None
|
mw._redis = None
|
||||||
m = make_message()
|
m = make_message()
|
||||||
key = mw._build_deduplication_key(m)
|
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):
|
def test_empty_kwargs_produces_consistent_key(self, middleware, make_message):
|
||||||
m1 = make_message(kwargs={})
|
m1 = make_message(kwargs={})
|
||||||
@@ -128,6 +128,10 @@ class TestDefaultBuildDeduplicationKey:
|
|||||||
)
|
)
|
||||||
assert middleware._build_deduplication_key(m) == "taskiq:deduplication:my-lock"
|
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:
|
class TestPreSend:
|
||||||
@pytest.mark.anyio
|
@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": 1}))
|
||||||
await middleware.pre_send(make_message(kwargs={"x": 2}))
|
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
|
@pytest.mark.anyio
|
||||||
async def test_default_ttl_applied(self, fake_redis, make_message):
|
async def test_default_ttl_applied(self, fake_redis, make_message):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost", default_ttl=77)
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost", default_ttl=77)
|
||||||
|
|||||||
Reference in New Issue
Block a user