feat: add RedisDeduplicationScheduleSource for scheduled task deduplication

This commit is contained in:
2026-07-16 21:37:19 +02:00
parent 002ebd2d47
commit 7a36993860
6 changed files with 391 additions and 33 deletions
+2
View File
@@ -1,10 +1,12 @@
"""Redis-backed deduplication middleware for Taskiq.""" """Redis-backed deduplication middleware for Taskiq."""
from .middleware import DuplicateTaskError, RedisDeduplicationMiddleware from .middleware import DuplicateTaskError, RedisDeduplicationMiddleware
from .schedule import RedisDeduplicationScheduleSource
__version__ = "1.1.0" __version__ = "1.1.0"
__all__ = [ __all__ = [
"DuplicateTaskError", "DuplicateTaskError",
"RedisDeduplicationMiddleware", "RedisDeduplicationMiddleware",
"RedisDeduplicationScheduleSource",
] ]
+46 -25
View File
@@ -137,29 +137,34 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
await self._redis.aclose() await self._redis.aclose()
def _build_deduplication_key(self, message: TaskiqMessage) -> str | None: def _build_deduplication_key(self, message: TaskiqMessage) -> str | None:
explicit_key: str | None = message.labels.get(DEDUP_EXPLICIT_KEY_LABEL) return self._build_key(message.task_name, message.labels, message.kwargs)
def _build_key(
self, task_name: str, labels: dict[str, Any], kwargs: dict[str, Any]
) -> str | None:
explicit_key: str | None = 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}"
key_fields = parse_list_label( key_fields = parse_list_label(
message.labels.get(DEDUP_KEY_FIELDS_LABEL), DEDUP_KEY_FIELDS_LABEL labels.get(DEDUP_KEY_FIELDS_LABEL), DEDUP_KEY_FIELDS_LABEL
) )
if key_fields is not None: if key_fields is not None:
missing = [field for field in key_fields if field not in message.kwargs] missing = [field for field in key_fields if field not in kwargs]
if missing: if missing:
logger.warning( logger.warning(
"Task %s requested deduplication_key_fields %r but they are " "Task %s requested deduplication_key_fields %r but they are "
"absent from kwargs; they are dropped from the fingerprint, which " "absent from kwargs; they are dropped from the fingerprint, which "
"may cause distinct calls to collide.", "may cause distinct calls to collide.",
message.task_name, task_name,
missing, missing,
) )
kwargs = {k: v for k, v in message.kwargs.items() if k in key_fields} filtered_kwargs = {k: v for k, v in kwargs.items() if k in key_fields}
else: else:
kwargs = message.kwargs filtered_kwargs = kwargs
try: try:
payload = json.dumps( payload = json.dumps(
{"task": message.task_name, "kwargs": kwargs}, {"task": task_name, "kwargs": filtered_kwargs},
sort_keys=True, sort_keys=True,
) )
except TypeError: except TypeError:
@@ -172,6 +177,31 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
labels.get(DEDUP_LABEL), self.default_deduplication, DEDUP_LABEL labels.get(DEDUP_LABEL), self.default_deduplication, DEDUP_LABEL
) )
def _require_redis(self) -> Redis:
if self._redis is None:
raise RuntimeError(
"RedisDeduplicationMiddleware.startup() was never called."
)
return self._redis
@staticmethod
def _decode_task_id(value: bytes | str) -> str:
return value.decode() if isinstance(value, bytes) else value
async def _peek(
self, task_name: str, labels: dict[str, Any], kwargs: dict[str, Any]
) -> tuple[str, str] | None:
if not self._is_enabled(labels):
return None
redis = self._require_redis()
key = self._build_key(task_name, labels, kwargs)
if key is None:
return None
holder_task_id = await redis.get(key)
if holder_task_id is None:
return None
return key, self._decode_task_id(holder_task_id)
def _get_ttl(self, labels: dict[str, Any]) -> int: def _get_ttl(self, labels: dict[str, Any]) -> int:
return parse_int_label( return parse_int_label(
labels.get(DEDUP_TTL_LABEL, self.default_ttl), labels.get(DEDUP_TTL_LABEL, self.default_ttl),
@@ -180,12 +210,9 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
) )
async def _release_if_owned(self, key: str, task_id: str) -> None: async def _release_if_owned(self, key: str, task_id: str) -> None:
if self._redis is None: redis = self._require_redis()
raise RuntimeError(
"RedisDeduplicationMiddleware.startup() was never called."
)
if self._release_script is None: if self._release_script is None:
self._release_script = self._redis.register_script(RELEASE_LUA_SCRIPT) self._release_script = redis.register_script(RELEASE_LUA_SCRIPT)
released = await check_and_delete(self._release_script, key, task_id) 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)
@@ -193,12 +220,9 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
logger.debug("Skipped release of lock %s: not owned by this task", key) logger.debug("Skipped release of lock %s: not owned by this task", key)
async def _refresh_if_owned(self, key: str, task_id: str, ttl: int) -> bool: async def _refresh_if_owned(self, key: str, task_id: str, ttl: int) -> bool:
if self._redis is None: redis = self._require_redis()
raise RuntimeError(
"RedisDeduplicationMiddleware.startup() was never called."
)
if self._refresh_script is None: if self._refresh_script is None:
self._refresh_script = self._redis.register_script(REFRESH_LUA_SCRIPT) self._refresh_script = redis.register_script(REFRESH_LUA_SCRIPT)
return await check_and_refresh(self._refresh_script, key, task_id, ttl) return await check_and_refresh(self._refresh_script, key, task_id, ttl)
@staticmethod @staticmethod
@@ -213,10 +237,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
if not self._is_enabled(message.labels): if not self._is_enabled(message.labels):
return message return message
if self._redis is None: redis = self._require_redis()
raise RuntimeError(
"RedisDeduplicationMiddleware.startup() was never called."
)
key = self._build_deduplication_key(message) key = self._build_deduplication_key(message)
self._cache_key(message, key) self._cache_key(message, key)
if key is None: if key is None:
@@ -229,11 +250,11 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
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)
acquired = await self._redis.set(key, message.task_id, ex=ttl, nx=True) acquired = await redis.set(key, message.task_id, ex=ttl, nx=True)
if not acquired: if not acquired:
holder_task_id = await self._redis.get(key) holder_task_id = await redis.get(key)
if isinstance(holder_task_id, bytes): if holder_task_id is not None:
holder_task_id = holder_task_id.decode() holder_task_id = self._decode_task_id(holder_task_id)
logger.warning( logger.warning(
"Duplicate task %s dropped (key=%s, holder_task_id=%s).", "Duplicate task %s dropped (key=%s, holder_task_id=%s).",
message.task_name, message.task_name,
+77
View File
@@ -0,0 +1,77 @@
import logging
from taskiq import ScheduledTask, ScheduleSource
from taskiq.exceptions import ScheduledTaskCancelledError
from taskiq.utils import maybe_awaitable
from .middleware import RedisDeduplicationMiddleware
logger = logging.getLogger(__name__)
class RedisDeduplicationScheduleSource(ScheduleSource):
"""Skips scheduled firings whose fingerprint is already locked.
Wraps a ``ScheduleSource`` and peeks the lock in ``pre_send``, raising
``ScheduledTaskCancelledError`` on a hit so the scheduler skips the firing
cleanly instead of raising ``DuplicateTaskError`` out of ``kiq()``. The
atomic acquire/release lifecycle stays owned by ``middleware``.
Attributes:
source: The wrapped ``ScheduleSource``.
middleware: The ``RedisDeduplicationMiddleware`` instance registered
on the broker. Must be the same instance, so both share one Redis
connection and configuration. Its ``startup()`` must have run
before ``pre_send()`` is invoked.
"""
def __init__(
self,
source: ScheduleSource,
middleware: RedisDeduplicationMiddleware,
) -> None:
self.source = source
self.middleware = middleware
async def startup(self) -> None:
await self.source.startup()
async def shutdown(self) -> None:
await self.source.shutdown()
async def get_schedules(self) -> list[ScheduledTask]:
return await self.source.get_schedules()
async def add_schedule(self, schedule: ScheduledTask) -> None:
await self.source.add_schedule(schedule)
async def delete_schedule(self, schedule_id: str) -> None:
await self.source.delete_schedule(schedule_id)
async def post_send(self, task: ScheduledTask) -> None:
await maybe_awaitable(self.source.post_send(task))
async def pre_send(self, task: ScheduledTask) -> None:
await maybe_awaitable(self.source.pre_send(task))
try:
held = await self.middleware._peek(task.task_name, task.labels, task.kwargs)
except RuntimeError:
logger.error(
"RedisDeduplicationMiddleware.startup() was never called; "
"cannot deduplicate scheduled task %s.",
task.task_name,
)
raise
if held is None:
return
key, holder_task_id = held
logger.warning(
"Duplicate scheduled task %s skipped before dispatch "
"(key=%s, holder_task_id=%s).",
task.task_name,
key,
holder_task_id,
)
raise ScheduledTaskCancelledError()
+31 -1
View File
@@ -1,7 +1,9 @@
import pytest import pytest
import fakeredis.aioredis import fakeredis.aioredis
from redis.asyncio import Redis from redis.asyncio import Redis
from taskiq import TaskiqMessage, TaskiqResult from taskiq import ScheduledTask, TaskiqMessage, TaskiqResult
from taskiq_deduplication import RedisDeduplicationMiddleware
@pytest.fixture @pytest.fixture
@@ -16,6 +18,13 @@ async def fake_redis():
await client.aclose() await client.aclose()
@pytest.fixture
def middleware(fake_redis):
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
mw._redis = fake_redis
return mw
@pytest.fixture @pytest.fixture
async def real_redis(): async def real_redis():
client = Redis.from_url("redis://localhost:6379/15") client = Redis.from_url("redis://localhost:6379/15")
@@ -57,3 +66,24 @@ def make_result():
) )
return _make return _make
@pytest.fixture
def make_scheduled_task():
def _make(
task_name="my_task",
schedule_id="schedule-1",
labels=None,
kwargs=None,
cron="* * * * *",
):
return ScheduledTask(
task_name=task_name,
schedule_id=schedule_id,
labels=labels or {},
args=[],
kwargs=kwargs or {},
cron=cron,
)
return _make
-7
View File
@@ -12,13 +12,6 @@ from taskiq_deduplication.middleware import (
) )
@pytest.fixture
def middleware(fake_redis):
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
mw._redis = fake_redis
return mw
class TestDefaultBuildDeduplicationKey: class TestDefaultBuildDeduplicationKey:
def test_same_kwargs_same_key(self, middleware, make_message): def test_same_kwargs_same_key(self, middleware, make_message):
m1 = make_message(kwargs={"a": 1, "b": 2}) m1 = make_message(kwargs={"a": 1, "b": 2})
+235
View File
@@ -0,0 +1,235 @@
import logging
import pytest
from taskiq import InMemoryBroker, ScheduleSource, TaskiqScheduler
from taskiq.exceptions import ScheduledTaskCancelledError
from taskiq_deduplication import RedisDeduplicationMiddleware
from taskiq_deduplication.middleware import (
DEDUP_EXPLICIT_KEY_LABEL,
DEDUP_KEY_FIELDS_LABEL,
DEDUP_LABEL,
)
from taskiq_deduplication.schedule import RedisDeduplicationScheduleSource
class FakeScheduleSource(ScheduleSource):
def __init__(self):
self.startup_called = False
self.shutdown_called = False
self.schedules_to_return = []
self.added = []
self.deleted = []
self.pre_send_calls = []
self.post_send_calls = []
self.pre_send_raises = None
async def startup(self):
self.startup_called = True
async def shutdown(self):
self.shutdown_called = True
async def get_schedules(self):
return self.schedules_to_return
async def add_schedule(self, schedule):
self.added.append(schedule)
async def delete_schedule(self, schedule_id):
self.deleted.append(schedule_id)
async def pre_send(self, task):
self.pre_send_calls.append(task)
if self.pre_send_raises is not None:
raise self.pre_send_raises
async def post_send(self, task):
self.post_send_calls.append(task)
@pytest.fixture
def fake_source():
return FakeScheduleSource()
@pytest.fixture
def wrapper(fake_source, middleware):
return RedisDeduplicationScheduleSource(fake_source, middleware)
class TestDelegation:
async def test_startup_delegates(self, wrapper, fake_source):
await wrapper.startup()
assert fake_source.startup_called
async def test_shutdown_delegates(self, wrapper, fake_source):
await wrapper.shutdown()
assert fake_source.shutdown_called
async def test_get_schedules_delegates(
self, wrapper, fake_source, make_scheduled_task
):
task = make_scheduled_task()
fake_source.schedules_to_return = [task]
assert await wrapper.get_schedules() == [task]
async def test_add_schedule_delegates(
self, wrapper, fake_source, make_scheduled_task
):
task = make_scheduled_task()
await wrapper.add_schedule(task)
assert fake_source.added == [task]
async def test_delete_schedule_delegates(self, wrapper, fake_source):
await wrapper.delete_schedule("some-id")
assert fake_source.deleted == ["some-id"]
async def test_post_send_delegates(self, wrapper, fake_source, make_scheduled_task):
task = make_scheduled_task()
await wrapper.post_send(task)
assert fake_source.post_send_calls == [task]
async def test_pre_send_delegates(self, wrapper, fake_source, make_scheduled_task):
task = make_scheduled_task()
await wrapper.pre_send(task)
assert fake_source.pre_send_calls == [task]
class TestPreSend:
async def test_no_lock_held_passes(self, wrapper, make_scheduled_task):
task = make_scheduled_task()
result = await wrapper.pre_send(task)
assert result is None
async def test_wrapped_source_cancellation_propagates(
self, wrapper, fake_source, make_scheduled_task
):
fake_source.pre_send_raises = ScheduledTaskCancelledError()
with pytest.raises(ScheduledTaskCancelledError):
await wrapper.pre_send(make_scheduled_task())
async def test_lock_held_raises_scheduled_task_cancelled_error(
self, wrapper, middleware, make_message, make_scheduled_task
):
await middleware.pre_send(make_message(task_name="my_task", kwargs={"a": 1}))
task = make_scheduled_task(task_name="my_task", kwargs={"a": 1})
with pytest.raises(ScheduledTaskCancelledError):
await wrapper.pre_send(task)
async def test_peek_does_not_acquire_or_mutate(
self, wrapper, middleware, fake_redis, make_scheduled_task
):
task = make_scheduled_task(task_name="my_task", kwargs={"a": 1})
await wrapper.pre_send(task)
key = middleware._build_key(task.task_name, task.labels, task.kwargs)
assert not await fake_redis.exists(key)
assert task.labels == {}
assert task.kwargs == {"a": 1}
async def test_peek_is_read_only_when_lock_held(
self, wrapper, middleware, fake_redis, make_message, make_scheduled_task
):
held_msg = make_message(task_name="my_task", task_id="holder", kwargs={"a": 1})
await middleware.pre_send(held_msg)
key = middleware._build_deduplication_key(held_msg)
ttl_before = await fake_redis.ttl(key)
holder_before = await fake_redis.get(key)
task = make_scheduled_task(task_name="my_task", kwargs={"a": 1})
with pytest.raises(ScheduledTaskCancelledError):
await wrapper.pre_send(task)
# The peek must not have re-set the key (TTL untouched) or changed
# its owner.
assert await fake_redis.get(key) == holder_before
assert await fake_redis.ttl(key) <= ttl_before
async def test_deduplication_disabled_label_bypasses_peek(
self, wrapper, middleware, make_message, make_scheduled_task
):
await middleware.pre_send(make_message(task_name="my_task", kwargs={"a": 1}))
task = make_scheduled_task(
task_name="my_task", kwargs={"a": 1}, labels={DEDUP_LABEL: False}
)
await wrapper.pre_send(task) # should not raise
async def test_deduplication_key_label_respected(
self, wrapper, middleware, make_message, make_scheduled_task
):
await middleware.pre_send(
make_message(kwargs={"a": 1}, labels={DEDUP_EXPLICIT_KEY_LABEL: "fixed"})
)
task = make_scheduled_task(
kwargs={"a": 999}, labels={DEDUP_EXPLICIT_KEY_LABEL: "fixed"}
)
with pytest.raises(ScheduledTaskCancelledError):
await wrapper.pre_send(task)
async def test_deduplication_key_fields_label_respected(
self, wrapper, middleware, make_message, make_scheduled_task
):
await middleware.pre_send(
make_message(
kwargs={"a": 1, "b": 2},
labels={DEDUP_KEY_FIELDS_LABEL: ["a"]},
)
)
task = make_scheduled_task(
kwargs={"a": 1, "b": 999},
labels={DEDUP_KEY_FIELDS_LABEL: ["a"]},
)
with pytest.raises(ScheduledTaskCancelledError):
await wrapper.pre_send(task)
async def test_non_serializable_kwargs_skips_peek_silently(
self, wrapper, make_scheduled_task, caplog
):
task = make_scheduled_task(kwargs={"dt": object()})
with caplog.at_level(logging.WARNING, logger="taskiq_deduplication.schedule"):
await wrapper.pre_send(task) # should not raise
assert not any("non-JSON-serializable" in r.message for r in caplog.records)
async def test_pre_send_without_middleware_startup_raises_runtime_error(
self, fake_source, make_scheduled_task
):
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
w = RedisDeduplicationScheduleSource(fake_source, mw)
with pytest.raises(RuntimeError, match="startup"):
await w.pre_send(make_scheduled_task())
async def test_pre_send_without_middleware_startup_logs_before_raising(
self, fake_source, make_scheduled_task, caplog
):
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
w = RedisDeduplicationScheduleSource(fake_source, mw)
with caplog.at_level(logging.ERROR, logger="taskiq_deduplication.schedule"):
with pytest.raises(RuntimeError):
await w.pre_send(make_scheduled_task())
assert any("startup" in r.message for r in caplog.records)
async def test_different_kwargs_both_pass(self, wrapper, make_scheduled_task):
await wrapper.pre_send(make_scheduled_task(kwargs={"x": 1}))
await wrapper.pre_send(make_scheduled_task(kwargs={"x": 2}))
class TestSchedulerIntegration:
async def test_second_firing_cancelled_without_uncaught_exception(
self, middleware, fake_source, make_message, make_scheduled_task
):
broker = InMemoryBroker().with_middlewares(middleware)
wrapper = RedisDeduplicationScheduleSource(fake_source, middleware)
scheduler = TaskiqScheduler(broker=broker, sources=[wrapper])
task_name = "my_task"
# Simulate a still-running first firing by holding the lock directly.
await middleware.pre_send(
make_message(task_name=task_name, task_id="holder", kwargs={})
)
scheduled = make_scheduled_task(task_name=task_name, kwargs={})
# Must be cancelled gracefully by on_ready()'s own except clause, not
# raise DuplicateTaskError out of kiq() into an uncaught exception.
await scheduler.on_ready(wrapper, scheduled)