From e45aa8f9776e5b153c3df903efbbd59a7dcff3f5 Mon Sep 17 00:00:00 2001 From: d3vyce <44915747+d3vyce@users.noreply.github.com> Date: Fri, 26 Jun 2026 20:59:23 +0200 Subject: [PATCH] feat: allow pydantic RedisDsn in addition to str for redis_url (#62) --- docs/usage.md | 2 +- pyproject.toml | 1 + src/taskiq_deduplication/middleware.py | 8 +++++--- tests/test_middleware.py | 12 ++++++++++++ uv.lock | 2 ++ 5 files changed, 21 insertions(+), 4 deletions(-) diff --git a/docs/usage.md b/docs/usage.md index e9262d7..268586a 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -17,7 +17,7 @@ broker = ListQueueBroker("redis://localhost:6379").with_middlewares( | Parameter | Type | Default | Description | |---|---|---|---| -| `redis_url` | `str` | — | Redis connection URL passed to `Redis.from_url`. | +| `redis_url` | `str \| RedisDsn` | — | Redis connection URL passed to `Redis.from_url`. Accepts a plain string or a pydantic [`RedisDsn`](https://docs.pydantic.dev/latest/api/networks/#pydantic.networks.RedisDsn). | | `default_deduplication` | `bool` | `True` | Whether deduplication is enabled for all tasks by default. Set `False` to opt-in per task instead of opting out. | | `default_ttl` | `int` | `300` | Default lock TTL in seconds. Overridden per task with the `deduplication_ttl` label. | | `key_prefix` | `str` | `"taskiq:deduplication"` | Prefix for all Redis lock keys. | diff --git a/pyproject.toml b/pyproject.toml index b28da19..334d854 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -27,6 +27,7 @@ classifiers = [ "Typing :: Typed", ] dependencies = [ + "pydantic>=2.0.0", "redis>=7.0.0", "taskiq>=0.12.0", ] diff --git a/src/taskiq_deduplication/middleware.py b/src/taskiq_deduplication/middleware.py index 7b90cd0..57e410a 100644 --- a/src/taskiq_deduplication/middleware.py +++ b/src/taskiq_deduplication/middleware.py @@ -4,6 +4,7 @@ import json import logging from typing import Any, Awaitable, cast +from pydantic import RedisDsn from redis.asyncio import Redis from taskiq import TaskiqMessage, TaskiqResult from taskiq.abc.middleware import TaskiqMiddleware @@ -39,7 +40,8 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): on completion or error. Attributes: - redis_url: Redis connection URL passed to ``Redis.from_url``. + redis_url: Redis connection URL (``str`` or ``RedisDsn``) passed to + ``Redis.from_url``. default_deduplication: Whether deduplication is enabled by default. default_ttl: Default lock TTL in seconds. key_prefix: Prefix for all Redis lock keys. @@ -47,7 +49,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): def __init__( self, - redis_url: str, + redis_url: str | RedisDsn, default_deduplication: bool = True, default_ttl: int = 300, key_prefix: str = "taskiq:deduplication", @@ -66,7 +68,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware): async def startup(self) -> None: last_error: BaseException | None = None for attempt in range(self.startup_retries): - client = Redis.from_url(self.redis_url) + client = Redis.from_url(str(self.redis_url)) try: await cast(Awaitable[bool], client.ping()) self._redis = client diff --git a/tests/test_middleware.py b/tests/test_middleware.py index 6fc0fb1..f746544 100644 --- a/tests/test_middleware.py +++ b/tests/test_middleware.py @@ -1,6 +1,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from pydantic import RedisDsn, TypeAdapter from taskiq_deduplication import DuplicateTaskError, RedisDeduplicationMiddleware from taskiq_deduplication.middleware import ( @@ -341,6 +342,17 @@ class TestLifecycle: mock_from_url.assert_called_once_with("redis://localhost") assert mw._redis is mock_client + async def test_startup_accepts_redis_dsn(self): + dsn = TypeAdapter(RedisDsn).validate_python("redis://localhost:6379/0") + mw = RedisDeduplicationMiddleware(redis_url=dsn) + 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:6379/0") + assert mw._redis is mock_client + async def test_shutdown_closes_redis_client(self): mw = RedisDeduplicationMiddleware(redis_url="redis://localhost") mock_client = AsyncMock() diff --git a/uv.lock b/uv.lock index 58ed97f..8fa0bb3 100644 --- a/uv.lock +++ b/uv.lock @@ -1488,6 +1488,7 @@ name = "taskiq-deduplication" version = "1.0.5" source = { editable = "." } dependencies = [ + { name = "pydantic" }, { name = "redis" }, { name = "taskiq" }, ] @@ -1520,6 +1521,7 @@ tests = [ [package.metadata] requires-dist = [ + { name = "pydantic", specifier = ">=2.0.0" }, { name = "redis", specifier = ">=7.0.0" }, { name = "taskiq", specifier = ">=0.12.0" }, ]