mirror of
https://github.com/d3vyce/taskiq-deduplication.git
synced 2026-08-04 19:14:07 +00:00
refactor: cleanup middleware helpers, remove type hacks, and cache Lua script registration (#32)
This commit is contained in:
@@ -8,7 +8,7 @@ from redis.asyncio import Redis
|
|||||||
from taskiq import TaskiqMessage, TaskiqResult
|
from taskiq import TaskiqMessage, TaskiqResult
|
||||||
from taskiq.abc.middleware import TaskiqMiddleware
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -55,6 +55,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
self.startup_retries = startup_retries
|
self.startup_retries = startup_retries
|
||||||
self.startup_retry_delay = startup_retry_delay
|
self.startup_retry_delay = startup_retry_delay
|
||||||
self._redis: Redis | None = None
|
self._redis: Redis | None = None
|
||||||
|
self._release_script: Any = None
|
||||||
|
|
||||||
async def startup(self) -> None:
|
async def startup(self) -> None:
|
||||||
last_error: BaseException | None = None
|
last_error: BaseException | None = None
|
||||||
@@ -62,6 +63,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
try:
|
try:
|
||||||
self._redis = Redis.from_url(self.redis_url)
|
self._redis = Redis.from_url(self.redis_url)
|
||||||
await cast(Awaitable[bool], self._redis.ping())
|
await cast(Awaitable[bool], self._redis.ping())
|
||||||
|
self._release_script = self._redis.register_script(RELEASE_LUA_SCRIPT)
|
||||||
return
|
return
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
last_error = exc
|
last_error = exc
|
||||||
@@ -133,7 +135,9 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
"RedisDeduplicationMiddleware.startup() was never called."
|
"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:
|
if released:
|
||||||
logger.debug("Released lock %s", key)
|
logger.debug("Released lock %s", key)
|
||||||
else:
|
else:
|
||||||
@@ -181,11 +185,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
logger.debug("Lock %s acquired for task %s", key, message.task_name)
|
logger.debug("Lock %s acquired for task %s", key, message.task_name)
|
||||||
return message
|
return message
|
||||||
|
|
||||||
async def post_execute(
|
async def _release_lock(self, message: TaskiqMessage) -> None:
|
||||||
self,
|
|
||||||
message: TaskiqMessage,
|
|
||||||
result: TaskiqResult,
|
|
||||||
) -> None:
|
|
||||||
if not self._is_enabled(message.labels):
|
if not self._is_enabled(message.labels):
|
||||||
return
|
return
|
||||||
key = self._get_cached_key(message)
|
key = self._get_cached_key(message)
|
||||||
@@ -193,15 +193,17 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
return
|
return
|
||||||
await self._release_if_owned(key, message.task_id)
|
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(
|
async def on_error(
|
||||||
self,
|
self,
|
||||||
message: TaskiqMessage,
|
message: TaskiqMessage,
|
||||||
result: TaskiqResult,
|
result: TaskiqResult,
|
||||||
exception: BaseException,
|
exception: BaseException,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not self._is_enabled(message.labels):
|
await self._release_lock(message)
|
||||||
return
|
|
||||||
key = self._get_cached_key(message)
|
|
||||||
if key is None:
|
|
||||||
return
|
|
||||||
await self._release_if_owned(key, message.task_id)
|
|
||||||
|
|||||||
@@ -1,5 +1,4 @@
|
|||||||
from redis.asyncio import Redis
|
from typing import Any
|
||||||
from redis.commands.core import AsyncScript
|
|
||||||
|
|
||||||
RELEASE_LUA_SCRIPT = """
|
RELEASE_LUA_SCRIPT = """
|
||||||
if redis.call('get', KEYS[1]) == ARGV[1] then
|
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*.
|
"""Delete *key* only if its value equals *owner*.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
redis: Async Redis client.
|
script: Pre-registered Lua script object (from ``Redis.register_script``).
|
||||||
key: Lock key to delete.
|
key: Lock key to delete.
|
||||||
owner: Expected value of the key (task_id).
|
owner: Expected value of the key (task_id).
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
True if the key was deleted, False otherwise.
|
True if the key was deleted, False otherwise.
|
||||||
"""
|
"""
|
||||||
script: AsyncScript = redis.register_script(RELEASE_LUA_SCRIPT)
|
|
||||||
released: int = await script(keys=[key], args=[owner])
|
released: int = await script(keys=[key], args=[owner])
|
||||||
return bool(released)
|
return bool(released)
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
@@ -309,6 +309,7 @@ class TestLifecycle:
|
|||||||
assert mw._redis is None
|
assert mw._redis is None
|
||||||
with patch("redis.asyncio.Redis.from_url") as mock_from_url:
|
with patch("redis.asyncio.Redis.from_url") as mock_from_url:
|
||||||
mock_client = AsyncMock()
|
mock_client = AsyncMock()
|
||||||
|
mock_client.register_script = MagicMock()
|
||||||
mock_from_url.return_value = mock_client
|
mock_from_url.return_value = mock_client
|
||||||
await mw.startup()
|
await mw.startup()
|
||||||
mock_from_url.assert_called_once_with("redis://localhost")
|
mock_from_url.assert_called_once_with("redis://localhost")
|
||||||
@@ -349,6 +350,7 @@ class TestStartupRetry:
|
|||||||
ConnectionError("fail"),
|
ConnectionError("fail"),
|
||||||
None,
|
None,
|
||||||
]
|
]
|
||||||
|
mock_client.register_script = MagicMock()
|
||||||
mock_from_url.return_value = mock_client
|
mock_from_url.return_value = mock_client
|
||||||
await mw.startup()
|
await mw.startup()
|
||||||
assert mw._redis is mock_client
|
assert mw._redis is mock_client
|
||||||
@@ -378,6 +380,7 @@ class TestStartupRetry:
|
|||||||
)
|
)
|
||||||
with patch("redis.asyncio.Redis.from_url") as mock_from_url:
|
with patch("redis.asyncio.Redis.from_url") as mock_from_url:
|
||||||
mock_client = AsyncMock()
|
mock_client = AsyncMock()
|
||||||
|
mock_client.register_script = MagicMock()
|
||||||
mock_from_url.return_value = mock_client
|
mock_from_url.return_value = mock_client
|
||||||
await mw.startup()
|
await mw.startup()
|
||||||
mock_client.ping.assert_called_once()
|
mock_client.ping.assert_called_once()
|
||||||
@@ -399,6 +402,7 @@ class TestStartupRetry:
|
|||||||
ConnectionError("fail"),
|
ConnectionError("fail"),
|
||||||
None,
|
None,
|
||||||
]
|
]
|
||||||
|
mock_client.register_script = MagicMock()
|
||||||
mock_from_url.return_value = mock_client
|
mock_from_url.return_value = mock_client
|
||||||
await mw.startup()
|
await mw.startup()
|
||||||
assert mock_sleep.call_count == 2
|
assert mock_sleep.call_count == 2
|
||||||
|
|||||||
Reference in New Issue
Block a user