diff --git a/src/taskiq_deduplication/__init__.py b/src/taskiq_deduplication/__init__.py index 36e01ea..06cbd1d 100644 --- a/src/taskiq_deduplication/__init__.py +++ b/src/taskiq_deduplication/__init__.py @@ -1,10 +1,12 @@ """Redis-backed deduplication middleware for Taskiq.""" from .middleware import DuplicateTaskError, RedisDeduplicationMiddleware +from .schedule import RedisDeduplicationScheduleSource __version__ = "1.1.0" __all__ = [ "DuplicateTaskError", "RedisDeduplicationMiddleware", + "RedisDeduplicationScheduleSource", ] diff --git a/src/taskiq_deduplication/middleware.py b/src/taskiq_deduplication/middleware.py index 8ad748f..244cc9d 100644 --- a/src/taskiq_deduplication/middleware.py +++ b/src/taskiq_deduplication/middleware.py @@ -137,29 +137,34 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): await self._redis.aclose() 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: return f"{self.key_prefix}:{explicit_key}" 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: - 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: logger.warning( "Task %s requested deduplication_key_fields %r but they are " "absent from kwargs; they are dropped from the fingerprint, which " "may cause distinct calls to collide.", - message.task_name, + task_name, 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: - kwargs = message.kwargs + filtered_kwargs = kwargs try: payload = json.dumps( - {"task": message.task_name, "kwargs": kwargs}, + {"task": task_name, "kwargs": filtered_kwargs}, sort_keys=True, ) except TypeError: @@ -172,6 +177,31 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): 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: return parse_int_label( 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: - if self._redis is None: - raise RuntimeError( - "RedisDeduplicationMiddleware.startup() was never called." - ) + redis = self._require_redis() 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) if released: 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) async def _refresh_if_owned(self, key: str, task_id: str, ttl: int) -> bool: - if self._redis is None: - raise RuntimeError( - "RedisDeduplicationMiddleware.startup() was never called." - ) + redis = self._require_redis() 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) @staticmethod @@ -213,10 +237,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): if not self._is_enabled(message.labels): return message - if self._redis is None: - raise RuntimeError( - "RedisDeduplicationMiddleware.startup() was never called." - ) + redis = self._require_redis() key = self._build_deduplication_key(message) self._cache_key(message, key) if key is None: @@ -229,11 +250,11 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): ttl = self._get_ttl(message.labels) 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: - holder_task_id = await self._redis.get(key) - if isinstance(holder_task_id, bytes): - holder_task_id = holder_task_id.decode() + holder_task_id = await redis.get(key) + if holder_task_id is not None: + holder_task_id = self._decode_task_id(holder_task_id) logger.warning( "Duplicate task %s dropped (key=%s, holder_task_id=%s).", message.task_name, diff --git a/src/taskiq_deduplication/schedule.py b/src/taskiq_deduplication/schedule.py new file mode 100644 index 0000000..43327c6 --- /dev/null +++ b/src/taskiq_deduplication/schedule.py @@ -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() diff --git a/tests/conftest.py b/tests/conftest.py index a78cae5..d83c2cd 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,7 +1,9 @@ import pytest import fakeredis.aioredis from redis.asyncio import Redis -from taskiq import TaskiqMessage, TaskiqResult +from taskiq import ScheduledTask, TaskiqMessage, TaskiqResult + +from taskiq_deduplication import RedisDeduplicationMiddleware @pytest.fixture @@ -16,6 +18,13 @@ async def fake_redis(): await client.aclose() +@pytest.fixture +def middleware(fake_redis): + mw = RedisDeduplicationMiddleware(redis_url="redis://localhost") + mw._redis = fake_redis + return mw + + @pytest.fixture async def real_redis(): client = Redis.from_url("redis://localhost:6379/15") @@ -57,3 +66,24 @@ def make_result(): ) 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 diff --git a/tests/test_middleware.py b/tests/test_middleware.py index c214149..ab881ce 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -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: def test_same_kwargs_same_key(self, middleware, make_message): m1 = make_message(kwargs={"a": 1, "b": 2}) diff --git a/tests/test_schedule.py b/tests/test_schedule.py new file mode 100644 index 0000000..d1a24e2 --- /dev/null +++ b/tests/test_schedule.py @@ -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)