Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
84 changes: 84 additions & 0 deletions tests/backends/test_redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -356,13 +357,15 @@ 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)
assert restored is not None
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()

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down
27 changes: 25 additions & 2 deletions threadmill/backends/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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 = {
Expand Down
5 changes: 5 additions & 0 deletions threadmill/backends/lua/acknowledge.lua
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
8 changes: 6 additions & 2 deletions threadmill/backends/lua/acquire.lua
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
29 changes: 23 additions & 6 deletions threadmill/backends/redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
TelemetryDirection,
TelemetryEvent,
ThreadmillTaskBackend,
ThreadmillTaskResult,
)
from threadmill.exceptions import AcknowledgementTimeout

Expand Down Expand Up @@ -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"

Expand All @@ -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:
Expand All @@ -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:
Expand Down Expand Up @@ -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),
Expand All @@ -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:
Expand All @@ -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,
Expand Down Expand Up @@ -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=[
Expand All @@ -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 "",
],
)

Expand Down Expand Up @@ -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())
Expand Down
Loading