diff --git a/README.md b/README.md index 893ef6f..128e84f 100644 --- a/README.md +++ b/README.md @@ -137,7 +137,9 @@ The `RedisTaskBackend` accepts the following options under `OPTIONS` in your A task whose lease expired reaches the `retry` callback as an `AcknowledgementTimeout` error, or is marked FAILED when nothing retries it. Keep `lease_ttl` above your worst-case runtime: a task that outlives its lease -can still be running, so a retry may execute concurrently with it. +can still be running, so a retry may execute concurrently with it. The +acknowledgement of the lease holder wins: the late result of an expired attempt +is discarded. All keys for one backend alias share a Redis Cluster hash tag (`{alias}`), so every multi-key operation — including the cross-queue acquire — runs on a single diff --git a/tests/backends/test_redis.py b/tests/backends/test_redis.py index 1a3f1f9..08a8284 100644 --- a/tests/backends/test_redis.py +++ b/tests/backends/test_redis.py @@ -13,6 +13,7 @@ import pytest from django.tasks import default_task_backend from django.tasks.base import TaskResultStatus +from django.tasks.exceptions import TaskResultDoesNotExist from django.utils import timezone from tests.testapp.tasks import ( @@ -356,6 +357,7 @@ def test_acquire__stamps_lease_on_task_hash(self): assert acquired.last_attempted_at is not None assert acquired.started_at == acquired.last_attempted_at assert acquired.worker_ids == ["test-worker"] + assert acquired.lease_token is not None # Verify the lease is persisted, not only applied in memory. restored = backend.get_leased_task(task_result.id) @@ -363,6 +365,7 @@ def test_acquire__stamps_lease_on_task_hash(self): assert restored.worker_ids == acquired.worker_ids assert restored.started_at == acquired.started_at assert restored.last_attempted_at == acquired.last_attempted_at + assert restored.lease_token == acquired.lease_token finally: backend.close() @@ -478,6 +481,31 @@ def test_acknowledge__stores_the_leased_attempt(self): finally: backend.close() + def test_acknowledge__keeps_lease_token_out_of_the_result(self): + """Persist no lease token: it is attempt state, not result state.""" + backend = _make_backend("acknowledge_token_test") + try: + task_result = backend.enqueue(echo, args=[42]) + acquired = backend.acquire( + timeout=datetime.timedelta(seconds=1), worker="token-test" + ) + assert acquired is not None + + backend.acknowledge( + replace( + acquired, + status=TaskResultStatus.SUCCESSFUL, + finished_at=timezone.now(), + ) + ) + + result_key = backend.RESULT_KEY.format( + prefix=backend.key_prefix, result_id=task_result.id + ) + assert "lease_token" not in json.loads(backend.client.get(result_key)) + finally: + backend.close() + async def test_running_reaper__fails_expired_tasks(self): """Running reaper creates FAILED results for tasks with expired lease.""" backend = RedisTaskBackend( @@ -565,6 +593,62 @@ def test_stale_acknowledge__is_noop(self): finally: backend.close() + def test_stale_acknowledge__keeps_retry_attempt_after_requeue(self): + """Discard a late acknowledgement while a retry attempt holds the lease.""" + backend = _make_backend( + "stale_ack_retry_test", lease_ttl=datetime.timedelta(seconds=1) + ) + try: + task_result = backend.enqueue(echo_retry_on_lease_expiry, args=[42]) + expired = backend.acquire( + timeout=datetime.timedelta(seconds=1), worker="expired-worker" + ) + assert expired is not None + _expire_lease(backend, task_result.id) + RedisBroker(backend).main() + + deferred_key = backend.DEFERRED_KEY.format( + prefix=backend.key_prefix, queue_name="default" + ) + backend.client.zadd(deferred_key, {task_result.id: 0}) + RedisBroker(backend).main() + retry = backend.acquire( + timeout=datetime.timedelta(seconds=1), worker="retry-worker" + ) + assert retry is not None + assert retry.id == task_result.id + + backend.acknowledge( + replace( + expired, + status=TaskResultStatus.SUCCESSFUL, + finished_at=timezone.now(), + ) + ) + + with pytest.raises(TaskResultDoesNotExist): + backend.get_result(task_result.id) + running_key = backend._segment_key(TaskResultStatus.RUNNING, "default") + task_key = backend.TASK_KEY.format( + prefix=backend.key_prefix, task_id=task_result.id + ) + assert backend.client.zscore(running_key, task_result.id) is not None + assert backend.client.exists(task_key) + + backend.acknowledge( + replace( + retry, + status=TaskResultStatus.SUCCESSFUL, + finished_at=timezone.now(), + ) + ) + assert backend.get_result(task_result.id).worker_ids == [ + "expired-worker", + "retry-worker", + ] + finally: + backend.close() + async def test_queue_stats__empty_backend(self): """queue_stats returns zero counts for an empty backend.""" backend = RedisTaskBackend( diff --git a/threadmill/backends/base.py b/threadmill/backends/base.py index 959a85f..8d48307 100644 --- a/threadmill/backends/base.py +++ b/threadmill/backends/base.py @@ -65,6 +65,24 @@ def __reduce__(self): return (reconstructor, (kwargs,)) +@dataclasses.dataclass(frozen=True, slots=True, kw_only=True) +class ThreadmillTaskResult(TaskResult): + lease_token: str | None = None + + @classmethod + def from_result( + cls, task_result: TaskResult, *, lease_token: str | None + ) -> ThreadmillTaskResult: + return cls( + **{ + field.name: getattr(task_result, field.name) + for field in dataclasses.fields(TaskResult) + if field.init + }, + lease_token=lease_token, + ) + + @dataclasses.dataclass(kw_only=True, slots=True) class QueueCounts: """Point-in-time cardinality of each queue segment.""" @@ -155,10 +173,15 @@ class TaskResultEncoder(DjangoJSONEncoder): """JSON encoder for TaskResult and TaskError objects.""" def default(self, o): - if isinstance(o, (TaskResult, TaskError)): + if isinstance(o, TaskResult): + return { + field.name: getattr(o, field.name) + for field in dataclasses.fields(TaskResult) + } + if isinstance(o, TaskError): return { field.name: getattr(o, field.name) - for field in dataclasses.fields(type(o)) + for field in dataclasses.fields(TaskError) } if isinstance(o, RetryTask): data = { diff --git a/threadmill/backends/lua/acknowledge.lua b/threadmill/backends/lua/acknowledge.lua index 896f02f..f112600 100644 --- a/threadmill/backends/lua/acknowledge.lua +++ b/threadmill/backends/lua/acknowledge.lua @@ -20,6 +20,11 @@ -- ARGV[7] -- queue name -- Returns: 1 on success, 0 if task was not in the running set +local lease_token = redis.call('HGET', KEYS[3], 'lease_token') +if lease_token and lease_token ~= ARGV[8] then + return 0 +end + local removed = redis.call('ZREM', KEYS[1], ARGV[1]) if removed == 0 then return 0 -- Task already reaped, skip diff --git a/threadmill/backends/lua/acquire.lua b/threadmill/backends/lua/acquire.lua index 5352e53..fda2c0b 100644 --- a/threadmill/backends/lua/acquire.lua +++ b/threadmill/backends/lua/acquire.lua @@ -31,9 +31,13 @@ for offset = 0, num_queues - 1 do local data = redis.call('HGET', task_key, 'data') if data then local deadline = tonumber(ARGV[1]) + lease_ttl_ms + local lease_token = string.format( + '%06x%06x%06x%06x', + math.random(0, 0xffffff), math.random(0, 0xffffff), + math.random(0, 0xffffff), math.random(0, 0xffffff)) redis.call('ZADD', KEYS[queue_index * 2 - 1], deadline, task_id) - redis.call('HSET', task_key, 'lease_worker', ARGV[5], 'lease_started_at', ARGV[2]) - return data + redis.call('HSET', task_key, 'lease_worker', ARGV[5], 'lease_started_at', ARGV[2], 'lease_token', lease_token) + return { data, lease_token } end end end diff --git a/threadmill/backends/redis.py b/threadmill/backends/redis.py index 1becfcc..f6a0630 100644 --- a/threadmill/backends/redis.py +++ b/threadmill/backends/redis.py @@ -27,6 +27,7 @@ TelemetryDirection, TelemetryEvent, ThreadmillTaskBackend, + ThreadmillTaskResult, ) from threadmill.exceptions import AcknowledgementTimeout @@ -166,7 +167,7 @@ class RedisTaskBackend(ThreadmillTaskBackend): SEGMENT_KEY = "{prefix}:{queue_name}:{status}" DEFERRED_KEY = "{prefix}:{queue_name}:deferred" - LEASE_FIELDS = ("lease_worker", "lease_started_at") + LEASE_FIELDS = ("lease_worker", "lease_started_at", "lease_token") TELEMETRY_CHANNEL = "{prefix}:telemetry" @@ -182,10 +183,10 @@ def _segment_key(self, status: TaskResultStatus, queue_name: str) -> str: status=status.value.lower(), ) - def get_leased_task(self, task_id: str) -> TaskResult | None: + def get_leased_task(self, task_id: str) -> ThreadmillTaskResult | None: """Return a running task as its lease holds it, or None when its hash is gone.""" task_key = self.TASK_KEY.format(prefix=self.key_prefix, task_id=task_id) - data, lease_worker, lease_started_at = self.client.hmget( + data, lease_worker, lease_started_at, lease_token = self.client.hmget( task_key, "data", *self.LEASE_FIELDS ) if data is None: @@ -194,6 +195,7 @@ def get_leased_task(self, task_id: str) -> TaskResult | None: self.deserialize_task_result(data.decode()), worker=lease_worker.decode() if lease_worker else None, lease_started_at=_parse_lease_started_at(lease_started_at), + lease_token=lease_token.decode() if lease_token is not None else None, ) def __init__(self, alias: str, params: dict) -> None: @@ -324,7 +326,7 @@ def acquire( now_ms = now.timestamp() * 1000 now_iso = now.isoformat() - if data := self._acquire_script( + if result := self._acquire_script( keys=keys, args=[ str(now_ms), @@ -336,12 +338,14 @@ def acquire( str(self._rotation_offset), ], ): + data, lease_token = result self._miss_count = 0 self._rotation_offset = (self._rotation_offset + 1) % len(queue_names) return self._apply_lease( self.deserialize_task_result(data.decode()), worker=worker, lease_started_at=now, + lease_token=lease_token.decode(), ) try: @@ -367,12 +371,16 @@ def _apply_lease( *, worker: str | None, lease_started_at: datetime.datetime | None, - ) -> TaskResult: + lease_token: str | None, + ) -> ThreadmillTaskResult: """Return a stored task result as a running attempt. `worker=None` means the lease records no worker; an empty string still counts as an attempt, and a task that already records a start keeps it. """ + task_result = ThreadmillTaskResult.from_result( + task_result, lease_token=lease_token + ) return dataclasses.replace( task_result, status=TaskResultStatus.RUNNING, @@ -402,6 +410,11 @@ def acknowledge(self, task_result: TaskResult) -> None: ) finished_at = task_result.finished_at or timezone.now() finish_score = finished_at.timestamp() * 1000 + lease_token = ( + task_result.lease_token + if isinstance(task_result, ThreadmillTaskResult) + else None + ) self._acknowledge_script( keys=[ @@ -419,6 +432,7 @@ def acknowledge(self, task_result: TaskResult) -> None: task_result.status.name, self.telemetry_channel, task_result.task.queue_name, + lease_token or "", ], ) @@ -518,13 +532,16 @@ def _peek_tasks( if not stored: continue if leased: - data, lease_worker, lease_started_at = stored + data, lease_worker, lease_started_at, lease_token = stored if not data: continue yield self._apply_lease( self.deserialize_task_result(data.decode()), worker=lease_worker.decode() if lease_worker else None, lease_started_at=_parse_lease_started_at(lease_started_at), + lease_token=( + lease_token.decode() if lease_token is not None else None + ), ) else: yield self.deserialize_task_result(stored.decode())