mirror of
https://github.com/d3vyce/taskiq-deduplication.git
synced 2026-08-04 19:14:07 +00:00
test: add real-Redis integration tests (#45)
This commit is contained in:
@@ -73,8 +73,8 @@ jobs:
|
|||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: uv sync --group dev
|
run: uv sync --group dev
|
||||||
|
|
||||||
- name: Run tests with coverage
|
- name: Run unit tests with coverage
|
||||||
run: uv run pytest --cov --cov-report=xml --cov-report=term-missing --junitxml=junit.xml -o junit_family=legacy
|
run: uv run pytest -m "not integration" --cov --cov-report=xml --cov-report=term-missing --junitxml=junit.xml -o junit_family=legacy
|
||||||
|
|
||||||
- name: Upload coverage to Codecov
|
- name: Upload coverage to Codecov
|
||||||
if: matrix.python-version == '3.14'
|
if: matrix.python-version == '3.14'
|
||||||
@@ -91,3 +91,39 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
token: ${{ secrets.CODECOV_TOKEN }}
|
token: ${{ secrets.CODECOV_TOKEN }}
|
||||||
report_type: test_results
|
report_type: test_results
|
||||||
|
|
||||||
|
test-integration:
|
||||||
|
name: Integration (${{ matrix.backend.name }})
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
strategy:
|
||||||
|
fail-fast: false
|
||||||
|
matrix:
|
||||||
|
backend:
|
||||||
|
- { name: redis, image: "redis:7" }
|
||||||
|
- { name: valkey, image: "valkey/valkey:8" }
|
||||||
|
|
||||||
|
services:
|
||||||
|
store:
|
||||||
|
image: ${{ matrix.backend.image }}
|
||||||
|
ports:
|
||||||
|
- 6379:6379
|
||||||
|
options: >-
|
||||||
|
--health-cmd "redis-cli ping"
|
||||||
|
--health-interval 5s
|
||||||
|
--health-timeout 3s
|
||||||
|
--health-retries 5
|
||||||
|
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v6
|
||||||
|
|
||||||
|
- name: Install uv
|
||||||
|
uses: astral-sh/setup-uv@v7
|
||||||
|
|
||||||
|
- name: Set up Python
|
||||||
|
run: uv python install 3.13
|
||||||
|
|
||||||
|
- name: Install dependencies
|
||||||
|
run: uv sync --group dev
|
||||||
|
|
||||||
|
- name: Run integration tests
|
||||||
|
run: uv run pytest -m integration -v
|
||||||
|
|||||||
@@ -64,9 +64,13 @@ build-backend = "uv_build"
|
|||||||
|
|
||||||
[tool.pytest.ini_options]
|
[tool.pytest.ini_options]
|
||||||
testpaths = ["tests"]
|
testpaths = ["tests"]
|
||||||
|
anyio_mode = "auto"
|
||||||
filterwarnings = [
|
filterwarnings = [
|
||||||
"ignore::DeprecationWarning",
|
"ignore::DeprecationWarning",
|
||||||
]
|
]
|
||||||
|
markers = [
|
||||||
|
"integration: requires a live Redis at localhost:6379",
|
||||||
|
]
|
||||||
|
|
||||||
[tool.coverage.run]
|
[tool.coverage.run]
|
||||||
source = ["src/taskiq_deduplication"]
|
source = ["src/taskiq_deduplication"]
|
||||||
|
|||||||
@@ -116,6 +116,19 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _parse_int_label(value: Any, default: int, label_name: str = "") -> int:
|
||||||
|
try:
|
||||||
|
return int(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
logger.warning(
|
||||||
|
"Invalid %r value %r (expected int); falling back to default (%d).",
|
||||||
|
label_name,
|
||||||
|
value,
|
||||||
|
default,
|
||||||
|
)
|
||||||
|
return default
|
||||||
|
|
||||||
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)
|
explicit_key: str | None = message.labels.get(DEDUP_EXPLICIT_KEY_LABEL)
|
||||||
if explicit_key is not None:
|
if explicit_key is not None:
|
||||||
@@ -145,16 +158,11 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _get_ttl(self, labels: dict[str, Any]) -> int:
|
def _get_ttl(self, labels: dict[str, Any]) -> int:
|
||||||
value = labels.get(DEDUP_TTL_LABEL, self.default_ttl)
|
return self._parse_int_label(
|
||||||
try:
|
labels.get(DEDUP_TTL_LABEL, self.default_ttl),
|
||||||
return int(value)
|
self.default_ttl,
|
||||||
except (TypeError, ValueError):
|
DEDUP_TTL_LABEL,
|
||||||
logger.warning(
|
)
|
||||||
"Invalid deduplication_ttl value %r; falling back to default (%ds).",
|
|
||||||
value,
|
|
||||||
self.default_ttl,
|
|
||||||
)
|
|
||||||
return self.default_ttl
|
|
||||||
|
|
||||||
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:
|
if self._redis is None:
|
||||||
@@ -212,6 +220,7 @@ class RedisDeduplicationMiddleware(TaskiqMiddleware):
|
|||||||
return message
|
return message
|
||||||
|
|
||||||
async def _release_lock(self, message: TaskiqMessage) -> None:
|
async def _release_lock(self, message: TaskiqMessage) -> None:
|
||||||
|
# The cached key is set by pre_send() only when deduplication is enabled.
|
||||||
key = self._get_cached_key(message)
|
key = self._get_cached_key(message)
|
||||||
if key is None:
|
if key is None:
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -1,5 +1,8 @@
|
|||||||
|
from typing import Awaitable, cast
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import fakeredis.aioredis
|
import fakeredis.aioredis
|
||||||
|
from redis.asyncio import Redis
|
||||||
from taskiq import TaskiqMessage, TaskiqResult
|
from taskiq import TaskiqMessage, TaskiqResult
|
||||||
|
|
||||||
|
|
||||||
@@ -15,6 +18,21 @@ async def fake_redis():
|
|||||||
await client.aclose()
|
await client.aclose()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def real_redis():
|
||||||
|
client = Redis.from_url("redis://localhost:6379/15")
|
||||||
|
try:
|
||||||
|
await cast(Awaitable[bool], client.ping())
|
||||||
|
except Exception:
|
||||||
|
await client.aclose()
|
||||||
|
pytest.skip("Redis not available at localhost:6379")
|
||||||
|
return
|
||||||
|
await client.flushdb()
|
||||||
|
yield client
|
||||||
|
await client.flushdb()
|
||||||
|
await client.aclose()
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def make_message():
|
def make_message():
|
||||||
def _make(task_name="my_task", task_id="task-1", labels=None, kwargs=None):
|
def _make(task_name="my_task", task_id="task-1", labels=None, kwargs=None):
|
||||||
|
|||||||
@@ -0,0 +1,88 @@
|
|||||||
|
"""Integration tests against a live Redis instance (localhost:6379/15).
|
||||||
|
|
||||||
|
These tests are skipped automatically when Redis is not reachable.
|
||||||
|
In CI a Redis service is started before the test step.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from taskiq_deduplication import DuplicateTaskError, RedisDeduplicationMiddleware
|
||||||
|
from taskiq_deduplication.utils import RELEASE_LUA_SCRIPT
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mw(real_redis):
|
||||||
|
middleware = RedisDeduplicationMiddleware(redis_url="redis://localhost:6379/15")
|
||||||
|
middleware._redis = real_redis
|
||||||
|
middleware._release_script = real_redis.register_script(RELEASE_LUA_SCRIPT)
|
||||||
|
return middleware
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_full_lifecycle(mw, real_redis, make_message, make_result):
|
||||||
|
msg = make_message()
|
||||||
|
await mw.pre_send(msg)
|
||||||
|
key = mw._build_deduplication_key(msg)
|
||||||
|
assert await real_redis.exists(key)
|
||||||
|
|
||||||
|
await mw.post_execute(msg, make_result())
|
||||||
|
assert not await real_redis.exists(key)
|
||||||
|
|
||||||
|
await mw.pre_send(make_message())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_duplicate_rejected(mw, make_message):
|
||||||
|
await mw.pre_send(make_message())
|
||||||
|
with pytest.raises(DuplicateTaskError):
|
||||||
|
await mw.pre_send(make_message())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_redispatch_after_on_error(mw, make_message, make_result):
|
||||||
|
msg = make_message()
|
||||||
|
await mw.pre_send(msg)
|
||||||
|
await mw.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
||||||
|
await mw.pre_send(make_message())
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_lua_only_owner_can_release(mw, real_redis, make_message):
|
||||||
|
owner_msg = make_message(task_id="owner")
|
||||||
|
key = mw._build_deduplication_key(owner_msg)
|
||||||
|
await real_redis.set(key, "owner", ex=60)
|
||||||
|
|
||||||
|
await mw._release_if_owned(key, "intruder")
|
||||||
|
assert await real_redis.exists(key)
|
||||||
|
|
||||||
|
await mw._release_if_owned(key, "owner")
|
||||||
|
assert not await real_redis.exists(key)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_ttl_is_applied(mw, real_redis, make_message):
|
||||||
|
msg = make_message()
|
||||||
|
await mw.pre_send(msg)
|
||||||
|
key = mw._build_deduplication_key(msg)
|
||||||
|
ttl = await real_redis.ttl(key)
|
||||||
|
assert 0 < ttl <= mw.default_ttl
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_explicit_key_end_to_end(mw, real_redis, make_message, make_result):
|
||||||
|
from taskiq_deduplication.middleware import DEDUP_EXPLICIT_KEY_LABEL
|
||||||
|
|
||||||
|
msg = make_message(labels={DEDUP_EXPLICIT_KEY_LABEL: "my-lock"})
|
||||||
|
await mw.pre_send(msg)
|
||||||
|
assert await real_redis.exists("taskiq:deduplication:my-lock")
|
||||||
|
|
||||||
|
with pytest.raises(DuplicateTaskError):
|
||||||
|
await mw.pre_send(
|
||||||
|
make_message(
|
||||||
|
kwargs={"different": "kwargs"},
|
||||||
|
labels={DEDUP_EXPLICIT_KEY_LABEL: "my-lock"},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
await mw.post_execute(msg, make_result())
|
||||||
|
assert not await real_redis.exists("taskiq:deduplication:my-lock")
|
||||||
+26
-44
@@ -134,27 +134,23 @@ class TestDefaultBuildDeduplicationKey:
|
|||||||
|
|
||||||
|
|
||||||
class TestPreSend:
|
class TestPreSend:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_first_send_passes(self, middleware, make_message):
|
async def test_first_send_passes(self, middleware, make_message):
|
||||||
msg = make_message()
|
msg = make_message()
|
||||||
result = await middleware.pre_send(msg)
|
result = await middleware.pre_send(msg)
|
||||||
assert result is msg
|
assert result is msg
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_duplicate_raises(self, middleware, make_message):
|
async def test_duplicate_raises(self, middleware, make_message):
|
||||||
msg = make_message()
|
msg = make_message()
|
||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
with pytest.raises(DuplicateTaskError):
|
with pytest.raises(DuplicateTaskError):
|
||||||
await middleware.pre_send(make_message())
|
await middleware.pre_send(make_message())
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_deduplication_disabled_label(self, middleware, make_message):
|
async def test_deduplication_disabled_label(self, middleware, make_message):
|
||||||
msg1 = make_message(labels={DEDUP_LABEL: False})
|
msg1 = make_message(labels={DEDUP_LABEL: False})
|
||||||
msg2 = make_message(labels={DEDUP_LABEL: False})
|
msg2 = make_message(labels={DEDUP_LABEL: False})
|
||||||
await middleware.pre_send(msg1)
|
await middleware.pre_send(msg1)
|
||||||
await middleware.pre_send(msg2) # should not raise
|
await middleware.pre_send(msg2) # should not raise
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_deduplication_disabled_by_default_init(
|
async def test_deduplication_disabled_by_default_init(
|
||||||
self, fake_redis, make_message
|
self, fake_redis, make_message
|
||||||
):
|
):
|
||||||
@@ -165,7 +161,6 @@ class TestPreSend:
|
|||||||
await mw.pre_send(make_message())
|
await mw.pre_send(make_message())
|
||||||
await mw.pre_send(make_message()) # should not raise
|
await mw.pre_send(make_message()) # should not raise
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_ttl_applied(self, middleware, fake_redis, make_message):
|
async def test_ttl_applied(self, middleware, fake_redis, make_message):
|
||||||
msg = make_message(labels={DEDUP_TTL_LABEL: 42})
|
msg = make_message(labels={DEDUP_TTL_LABEL: 42})
|
||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
@@ -173,12 +168,10 @@ class TestPreSend:
|
|||||||
ttl = await fake_redis.ttl(key)
|
ttl = await fake_redis.ttl(key)
|
||||||
assert 0 < ttl <= 42
|
assert 0 < ttl <= 42
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_different_kwargs_both_pass(self, middleware, make_message):
|
async def test_different_kwargs_both_pass(self, middleware, make_message):
|
||||||
await middleware.pre_send(make_message(kwargs={"x": 1}))
|
await middleware.pre_send(make_message(kwargs={"x": 1}))
|
||||||
await middleware.pre_send(make_message(kwargs={"x": 2}))
|
await middleware.pre_send(make_message(kwargs={"x": 2}))
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_non_serializable_kwargs_skips_deduplication(
|
async def test_non_serializable_kwargs_skips_deduplication(
|
||||||
self, middleware, make_message
|
self, middleware, make_message
|
||||||
):
|
):
|
||||||
@@ -187,7 +180,6 @@ class TestPreSend:
|
|||||||
await middleware.pre_send(msg1)
|
await middleware.pre_send(msg1)
|
||||||
await middleware.pre_send(msg2) # should not raise
|
await middleware.pre_send(msg2) # should not raise
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_default_ttl_applied(self, fake_redis, make_message):
|
async def test_default_ttl_applied(self, fake_redis, make_message):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost", default_ttl=77)
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost", default_ttl=77)
|
||||||
mw._redis = fake_redis
|
mw._redis = fake_redis
|
||||||
@@ -199,7 +191,6 @@ class TestPreSend:
|
|||||||
|
|
||||||
|
|
||||||
class TestPostExecute:
|
class TestPostExecute:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_releases_lock(
|
async def test_releases_lock(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -211,7 +202,6 @@ class TestPostExecute:
|
|||||||
await middleware.post_execute(msg, make_result())
|
await middleware.post_execute(msg, make_result())
|
||||||
assert not await fake_redis.exists(key)
|
assert not await fake_redis.exists(key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_deduplication_disabled_noop(
|
async def test_deduplication_disabled_noop(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -223,7 +213,6 @@ class TestPostExecute:
|
|||||||
await middleware.post_execute(disabled_msg, make_result())
|
await middleware.post_execute(disabled_msg, make_result())
|
||||||
assert await fake_redis.exists(key)
|
assert await fake_redis.exists(key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_post_execute_after_ttl_expiry_is_safe(
|
async def test_post_execute_after_ttl_expiry_is_safe(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -234,7 +223,6 @@ class TestPostExecute:
|
|||||||
|
|
||||||
|
|
||||||
class TestOnError:
|
class TestOnError:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_releases_lock_on_error(
|
async def test_releases_lock_on_error(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -246,7 +234,6 @@ class TestOnError:
|
|||||||
await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
||||||
assert not await fake_redis.exists(key)
|
assert not await fake_redis.exists(key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_deduplication_disabled_noop(
|
async def test_deduplication_disabled_noop(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -262,7 +249,6 @@ class TestOnError:
|
|||||||
|
|
||||||
|
|
||||||
class TestSenderWorkerConfigMismatch:
|
class TestSenderWorkerConfigMismatch:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_post_execute_releases_lock_despite_worker_disabled_default(
|
async def test_post_execute_releases_lock_despite_worker_disabled_default(
|
||||||
self, fake_redis, make_message, make_result
|
self, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -284,7 +270,6 @@ class TestSenderWorkerConfigMismatch:
|
|||||||
await worker_mw.post_execute(msg, make_result())
|
await worker_mw.post_execute(msg, make_result())
|
||||||
assert not await fake_redis.exists(key)
|
assert not await fake_redis.exists(key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_on_error_releases_lock_despite_worker_disabled_default(
|
async def test_on_error_releases_lock_despite_worker_disabled_default(
|
||||||
self, fake_redis, make_message, make_result
|
self, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -308,7 +293,6 @@ class TestSenderWorkerConfigMismatch:
|
|||||||
|
|
||||||
|
|
||||||
class TestRedispatchAfterRelease:
|
class TestRedispatchAfterRelease:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_redispatch_after_post_execute(
|
async def test_redispatch_after_post_execute(
|
||||||
self, middleware, make_message, make_result
|
self, middleware, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -317,7 +301,6 @@ class TestRedispatchAfterRelease:
|
|||||||
await middleware.post_execute(msg, make_result())
|
await middleware.post_execute(msg, make_result())
|
||||||
await middleware.pre_send(make_message())
|
await middleware.pre_send(make_message())
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_redispatch_after_on_error(
|
async def test_redispatch_after_on_error(
|
||||||
self, middleware, make_message, make_result
|
self, middleware, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -328,7 +311,6 @@ class TestRedispatchAfterRelease:
|
|||||||
|
|
||||||
|
|
||||||
class TestAtomicRelease:
|
class TestAtomicRelease:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_only_owner_can_release(self, middleware, fake_redis, make_message):
|
async def test_only_owner_can_release(self, middleware, fake_redis, make_message):
|
||||||
owner_msg = make_message(task_id="owner-task")
|
owner_msg = make_message(task_id="owner-task")
|
||||||
key = middleware._build_deduplication_key(owner_msg)
|
key = middleware._build_deduplication_key(owner_msg)
|
||||||
@@ -341,7 +323,6 @@ class TestAtomicRelease:
|
|||||||
await middleware._release_if_owned(key, "owner-task")
|
await middleware._release_if_owned(key, "owner-task")
|
||||||
assert not await fake_redis.exists(key)
|
assert not await fake_redis.exists(key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_release_missing_key_is_noop(self, middleware, fake_redis):
|
async def test_release_missing_key_is_noop(self, middleware, fake_redis):
|
||||||
await middleware._release_if_owned(
|
await middleware._release_if_owned(
|
||||||
"taskiq:deduplication:nonexistent", "some-task"
|
"taskiq:deduplication:nonexistent", "some-task"
|
||||||
@@ -349,7 +330,6 @@ class TestAtomicRelease:
|
|||||||
|
|
||||||
|
|
||||||
class TestLifecycle:
|
class TestLifecycle:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_creates_redis_client(self):
|
async def test_startup_creates_redis_client(self):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
||||||
assert mw._redis is None
|
assert mw._redis is None
|
||||||
@@ -361,7 +341,6 @@ class TestLifecycle:
|
|||||||
mock_from_url.assert_called_once_with("redis://localhost")
|
mock_from_url.assert_called_once_with("redis://localhost")
|
||||||
assert mw._redis is mock_client
|
assert mw._redis is mock_client
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_shutdown_closes_redis_client(self):
|
async def test_shutdown_closes_redis_client(self):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
||||||
mock_client = AsyncMock()
|
mock_client = AsyncMock()
|
||||||
@@ -369,12 +348,10 @@ class TestLifecycle:
|
|||||||
await mw.shutdown()
|
await mw.shutdown()
|
||||||
mock_client.aclose.assert_called_once()
|
mock_client.aclose.assert_called_once()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_shutdown_without_startup_is_safe(self):
|
async def test_shutdown_without_startup_is_safe(self):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
||||||
await mw.shutdown()
|
await mw.shutdown()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_pre_send_without_startup_raises_runtime_error(self, make_message):
|
async def test_pre_send_without_startup_raises_runtime_error(self, make_message):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
||||||
with pytest.raises(RuntimeError, match="startup"):
|
with pytest.raises(RuntimeError, match="startup"):
|
||||||
@@ -382,7 +359,6 @@ class TestLifecycle:
|
|||||||
|
|
||||||
|
|
||||||
class TestStartupRetry:
|
class TestStartupRetry:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_succeeds_after_retries(self):
|
async def test_startup_succeeds_after_retries(self):
|
||||||
mw = RedisDeduplicationMiddleware(
|
mw = RedisDeduplicationMiddleware(
|
||||||
redis_url="redis://localhost",
|
redis_url="redis://localhost",
|
||||||
@@ -402,7 +378,6 @@ class TestStartupRetry:
|
|||||||
assert mw._redis is mock_client
|
assert mw._redis is mock_client
|
||||||
assert mock_client.ping.call_count == 3
|
assert mock_client.ping.call_count == 3
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_raises_after_all_retries_exhausted(self):
|
async def test_startup_raises_after_all_retries_exhausted(self):
|
||||||
mw = RedisDeduplicationMiddleware(
|
mw = RedisDeduplicationMiddleware(
|
||||||
redis_url="redis://localhost",
|
redis_url="redis://localhost",
|
||||||
@@ -417,7 +392,6 @@ class TestStartupRetry:
|
|||||||
await mw.startup()
|
await mw.startup()
|
||||||
assert mock_client.ping.call_count == 2
|
assert mock_client.ping.call_count == 2
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_no_retry_on_first_success(self):
|
async def test_startup_no_retry_on_first_success(self):
|
||||||
mw = RedisDeduplicationMiddleware(
|
mw = RedisDeduplicationMiddleware(
|
||||||
redis_url="redis://localhost",
|
redis_url="redis://localhost",
|
||||||
@@ -431,7 +405,6 @@ class TestStartupRetry:
|
|||||||
await mw.startup()
|
await mw.startup()
|
||||||
mock_client.ping.assert_called_once()
|
mock_client.ping.assert_called_once()
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_retry_delay_exponential(self):
|
async def test_startup_retry_delay_exponential(self):
|
||||||
mw = RedisDeduplicationMiddleware(
|
mw = RedisDeduplicationMiddleware(
|
||||||
redis_url="redis://localhost",
|
redis_url="redis://localhost",
|
||||||
@@ -455,7 +428,6 @@ class TestStartupRetry:
|
|||||||
assert mock_sleep.call_args_list[0].args[0] == 0.01
|
assert mock_sleep.call_args_list[0].args[0] == 0.01
|
||||||
assert mock_sleep.call_args_list[1].args[0] == 0.02
|
assert mock_sleep.call_args_list[1].args[0] == 0.02
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_closes_failed_clients(self):
|
async def test_startup_closes_failed_clients(self):
|
||||||
mw = RedisDeduplicationMiddleware(
|
mw = RedisDeduplicationMiddleware(
|
||||||
redis_url="redis://localhost",
|
redis_url="redis://localhost",
|
||||||
@@ -477,7 +449,6 @@ class TestStartupRetry:
|
|||||||
good.aclose.assert_not_called()
|
good.aclose.assert_not_called()
|
||||||
assert mw._redis is good
|
assert mw._redis is good
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_startup_closes_client_when_all_retries_exhausted(self):
|
async def test_startup_closes_client_when_all_retries_exhausted(self):
|
||||||
mw = RedisDeduplicationMiddleware(
|
mw = RedisDeduplicationMiddleware(
|
||||||
redis_url="redis://localhost",
|
redis_url="redis://localhost",
|
||||||
@@ -497,21 +468,18 @@ class TestStartupRetry:
|
|||||||
|
|
||||||
|
|
||||||
class TestLabelTypeParsing:
|
class TestLabelTypeParsing:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_bool_label_false_disables_dedup(self, fake_redis, make_message):
|
async def test_bool_label_false_disables_dedup(self, fake_redis, make_message):
|
||||||
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
mw = RedisDeduplicationMiddleware(redis_url="redis://localhost")
|
||||||
mw._redis = fake_redis
|
mw._redis = fake_redis
|
||||||
await mw.pre_send(make_message(labels={DEDUP_LABEL: False}))
|
await mw.pre_send(make_message(labels={DEDUP_LABEL: False}))
|
||||||
await mw.pre_send(make_message(labels={DEDUP_LABEL: False}))
|
await mw.pre_send(make_message(labels={DEDUP_LABEL: False}))
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_bool_label_true_enables_dedup(self, middleware, make_message):
|
async def test_bool_label_true_enables_dedup(self, middleware, make_message):
|
||||||
msg = make_message(labels={DEDUP_LABEL: True})
|
msg = make_message(labels={DEDUP_LABEL: True})
|
||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
with pytest.raises(DuplicateTaskError):
|
with pytest.raises(DuplicateTaskError):
|
||||||
await middleware.pre_send(make_message(labels={DEDUP_LABEL: True}))
|
await middleware.pre_send(make_message(labels={DEDUP_LABEL: True}))
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_key_fields_list_parsed_correctly(self, middleware, make_message):
|
async def test_key_fields_list_parsed_correctly(self, middleware, make_message):
|
||||||
m1 = make_message(
|
m1 = make_message(
|
||||||
kwargs={"a": 1, "b": 2, "c": 3},
|
kwargs={"a": 1, "b": 2, "c": 3},
|
||||||
@@ -525,7 +493,6 @@ class TestLabelTypeParsing:
|
|||||||
m1
|
m1
|
||||||
) == middleware._build_deduplication_key(m2)
|
) == middleware._build_deduplication_key(m2)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_key_fields_non_list_ignored(self, middleware, make_message):
|
async def test_key_fields_non_list_ignored(self, middleware, make_message):
|
||||||
m = make_message(
|
m = make_message(
|
||||||
kwargs={"a": 1},
|
kwargs={"a": 1},
|
||||||
@@ -534,7 +501,6 @@ class TestLabelTypeParsing:
|
|||||||
key = middleware._build_deduplication_key(m)
|
key = middleware._build_deduplication_key(m)
|
||||||
assert key is not None
|
assert key is not None
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_invalid_bool_label_warns_and_uses_default(
|
async def test_invalid_bool_label_warns_and_uses_default(
|
||||||
self, middleware, make_message, caplog
|
self, middleware, make_message, caplog
|
||||||
):
|
):
|
||||||
@@ -545,7 +511,6 @@ class TestLabelTypeParsing:
|
|||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
assert any("yes" in r.message for r in caplog.records)
|
assert any("yes" in r.message for r in caplog.records)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_invalid_key_fields_label_warns_and_falls_back_to_all_kwargs(
|
async def test_invalid_key_fields_label_warns_and_falls_back_to_all_kwargs(
|
||||||
self, middleware, make_message, caplog
|
self, middleware, make_message, caplog
|
||||||
):
|
):
|
||||||
@@ -558,7 +523,6 @@ class TestLabelTypeParsing:
|
|||||||
# falls back to full-kwargs fingerprint — key must still be produced
|
# falls back to full-kwargs fingerprint — key must still be produced
|
||||||
assert key is not None
|
assert key is not None
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_invalid_ttl_string_warns_and_uses_default(
|
async def test_invalid_ttl_string_warns_and_uses_default(
|
||||||
self, middleware, fake_redis, make_message, caplog
|
self, middleware, fake_redis, make_message, caplog
|
||||||
):
|
):
|
||||||
@@ -572,7 +536,6 @@ class TestLabelTypeParsing:
|
|||||||
ttl = await fake_redis.ttl(key)
|
ttl = await fake_redis.ttl(key)
|
||||||
assert 0 < ttl <= middleware.default_ttl
|
assert 0 < ttl <= middleware.default_ttl
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_invalid_ttl_none_warns_and_uses_default(
|
async def test_invalid_ttl_none_warns_and_uses_default(
|
||||||
self, middleware, fake_redis, make_message, caplog
|
self, middleware, fake_redis, make_message, caplog
|
||||||
):
|
):
|
||||||
@@ -588,7 +551,6 @@ class TestLabelTypeParsing:
|
|||||||
|
|
||||||
|
|
||||||
class TestKeyCaching:
|
class TestKeyCaching:
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_key_cached_during_pre_send(self, middleware, make_message):
|
async def test_key_cached_during_pre_send(self, middleware, make_message):
|
||||||
from taskiq_deduplication.middleware import _CACHED_KEY_LABEL
|
from taskiq_deduplication.middleware import _CACHED_KEY_LABEL
|
||||||
|
|
||||||
@@ -596,7 +558,6 @@ class TestKeyCaching:
|
|||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
assert _CACHED_KEY_LABEL in msg.labels
|
assert _CACHED_KEY_LABEL in msg.labels
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_post_execute_uses_cached_key(
|
async def test_post_execute_uses_cached_key(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -608,7 +569,6 @@ class TestKeyCaching:
|
|||||||
await middleware.post_execute(msg, make_result())
|
await middleware.post_execute(msg, make_result())
|
||||||
assert not await fake_redis.exists(cached_key)
|
assert not await fake_redis.exists(cached_key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_on_error_uses_cached_key(
|
async def test_on_error_uses_cached_key(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -620,7 +580,6 @@ class TestKeyCaching:
|
|||||||
await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
||||||
assert not await fake_redis.exists(cached_key)
|
assert not await fake_redis.exists(cached_key)
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_cached_key_none_when_key_build_fails(self, middleware, make_message):
|
async def test_cached_key_none_when_key_build_fails(self, middleware, make_message):
|
||||||
from taskiq_deduplication.middleware import _CACHED_KEY_LABEL
|
from taskiq_deduplication.middleware import _CACHED_KEY_LABEL
|
||||||
|
|
||||||
@@ -628,7 +587,6 @@ class TestKeyCaching:
|
|||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
assert msg.labels[_CACHED_KEY_LABEL] is None
|
assert msg.labels[_CACHED_KEY_LABEL] is None
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_post_execute_noop_when_cached_key_is_none(
|
async def test_post_execute_noop_when_cached_key_is_none(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -636,7 +594,6 @@ class TestKeyCaching:
|
|||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
await middleware.post_execute(msg, make_result())
|
await middleware.post_execute(msg, make_result())
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_on_error_noop_when_cached_key_is_none(
|
async def test_on_error_noop_when_cached_key_is_none(
|
||||||
self, middleware, fake_redis, make_message, make_result
|
self, middleware, fake_redis, make_message, make_result
|
||||||
):
|
):
|
||||||
@@ -644,10 +601,35 @@ class TestKeyCaching:
|
|||||||
await middleware.pre_send(msg)
|
await middleware.pre_send(msg)
|
||||||
await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
await middleware.on_error(msg, make_result(is_err=True), RuntimeError("boom"))
|
||||||
|
|
||||||
@pytest.mark.anyio
|
|
||||||
async def test_release_if_owned_raises_without_redis(
|
async def test_release_if_owned_raises_without_redis(
|
||||||
self, middleware, make_message
|
self, middleware, make_message
|
||||||
):
|
):
|
||||||
middleware._redis = None
|
middleware._redis = None
|
||||||
with pytest.raises(RuntimeError, match="startup"):
|
with pytest.raises(RuntimeError, match="startup"):
|
||||||
await middleware._release_if_owned("some-key", "some-task")
|
await middleware._release_if_owned("some-key", "some-task")
|
||||||
|
|
||||||
|
|
||||||
|
class TestTTLExpiry:
|
||||||
|
async def test_lock_expiry_admits_duplicate(
|
||||||
|
self, middleware, fake_redis, make_message
|
||||||
|
):
|
||||||
|
msg = make_message()
|
||||||
|
await middleware.pre_send(msg)
|
||||||
|
# simulate TTL expiry by deleting the key
|
||||||
|
await fake_redis.delete(middleware._build_deduplication_key(msg))
|
||||||
|
# same fingerprint should now pass since lock is gone
|
||||||
|
await middleware.pre_send(make_message())
|
||||||
|
|
||||||
|
|
||||||
|
class TestExplicitKeyEdgeCases:
|
||||||
|
def test_empty_string_key_produces_prefix_only_key(self, middleware, make_message):
|
||||||
|
m = make_message(labels={DEDUP_EXPLICIT_KEY_LABEL: ""})
|
||||||
|
key = middleware._build_deduplication_key(m)
|
||||||
|
assert key == "taskiq:deduplication:"
|
||||||
|
|
||||||
|
async def test_empty_string_key_acquires_lock(
|
||||||
|
self, middleware, fake_redis, make_message
|
||||||
|
):
|
||||||
|
msg = make_message(labels={DEDUP_EXPLICIT_KEY_LABEL: ""})
|
||||||
|
await middleware.pre_send(msg)
|
||||||
|
assert await fake_redis.exists("taskiq:deduplication:")
|
||||||
|
|||||||
Reference in New Issue
Block a user