refactor: cleanup middleware helpers, remove type hacks, and cache Lua script registration (#32)

This commit is contained in:
d3vyce
2026-05-15 21:52:39 +02:00
committed by GitHub
parent ac80781018
commit 0bc6a18d43
3 changed files with 23 additions and 19 deletions
+15 -13
View File
@@ -8,7 +8,7 @@ from redis.asyncio import Redis
from taskiq import TaskiqMessage, TaskiqResult
from taskiq.abc.middleware import TaskiqMiddleware
from .utils import check_and_delete
from .utils import RELEASE_LUA_SCRIPT, check_and_delete
logger = logging.getLogger(__name__)
@@ -55,6 +55,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
self.startup_retries = startup_retries
self.startup_retry_delay = startup_retry_delay
self._redis: Redis | None = None
self._release_script: Any = None
async def startup(self) -> None:
last_error: BaseException | None = None
@@ -62,6 +63,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
try:
self._redis = Redis.from_url(self.redis_url)
await cast(Awaitable[bool], self._redis.ping())
self._release_script = self._redis.register_script(RELEASE_LUA_SCRIPT)
return
except Exception as exc:
last_error = exc
@@ -133,7 +135,9 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
raise RuntimeError(
"RedisDeduplicationMiddleware.startup() was never called."
)
released = await check_and_delete(self._redis, key, task_id)
if self._release_script is None:
self._release_script = self._redis.register_script(RELEASE_LUA_SCRIPT)
released = await check_and_delete(self._release_script, key, task_id)
if released:
logger.debug("Released lock %s", key)
else:
@@ -181,11 +185,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
logger.debug("Lock %s acquired for task %s", key, message.task_name)
return message
async def post_execute(
self,
message: TaskiqMessage,
result: TaskiqResult,
) -> None:
async def _release_lock(self, message: TaskiqMessage) -> None:
if not self._is_enabled(message.labels):
return
key = self._get_cached_key(message)
@@ -193,15 +193,17 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
return
await self._release_if_owned(key, message.task_id)
async def post_execute(
self,
message: TaskiqMessage,
result: TaskiqResult,
) -> None:
await self._release_lock(message)
async def on_error(
self,
message: TaskiqMessage,
result: TaskiqResult,
exception: BaseException,
) -> None:
if not self._is_enabled(message.labels):
return
key = self._get_cached_key(message)
if key is None:
return
await self._release_if_owned(key, message.task_id)
await self._release_lock(message)
+3 -5
View File
@@ -1,5 +1,4 @@
from redis.asyncio import Redis
from redis.commands.core import AsyncScript
from typing import Any
RELEASE_LUA_SCRIPT = """
if redis.call('get', KEYS[1]) == ARGV[1] then
@@ -10,17 +9,16 @@ end
"""
async def check_and_delete(redis: Redis, key: str, owner: str) -> bool:
async def check_and_delete(script: Any, key: str, owner: str) -> bool:
"""Delete *key* only if its value equals *owner*.
Args:
redis: Async Redis client.
script: Pre-registered Lua script object (from ``Redis.register_script``).
key: Lock key to delete.
owner: Expected value of the key (task_id).
Returns:
True if the key was deleted, False otherwise.
"""
script: AsyncScript = redis.register_script(RELEASE_LUA_SCRIPT)
released: int = await script(keys=[key], args=[owner])
return bool(released)
+5 -1
View File
@@ -1,4 +1,4 @@
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -309,6 +309,7 @@ class TestLifecycle:
assert mw._redis is None
with patch("redis.asyncio.Redis.from_url") as mock_from_url:
mock_client = AsyncMock()
mock_client.register_script = MagicMock()
mock_from_url.return_value = mock_client
await mw.startup()
mock_from_url.assert_called_once_with("redis://localhost")
@@ -349,6 +350,7 @@ class TestStartupRetry:
ConnectionError("fail"),
None,
]
mock_client.register_script = MagicMock()
mock_from_url.return_value = mock_client
await mw.startup()
assert mw._redis is mock_client
@@ -378,6 +380,7 @@ class TestStartupRetry:
)
with patch("redis.asyncio.Redis.from_url") as mock_from_url:
mock_client = AsyncMock()
mock_client.register_script = MagicMock()
mock_from_url.return_value = mock_client
await mw.startup()
mock_client.ping.assert_called_once()
@@ -399,6 +402,7 @@ class TestStartupRetry:
ConnectionError("fail"),
None,
]
mock_client.register_script = MagicMock()
mock_from_url.return_value = mock_client
await mw.startup()
assert mock_sleep.call_count == 2