Skip to content

Commit 6cc4da6

Browse files
[feat] priority-based scheduling
1 parent 5ee4b48 commit 6cc4da6

8 files changed

Lines changed: 516 additions & 70 deletions

File tree

arq/connections.py

Lines changed: 87 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,18 @@
1313
from redis.asyncio.sentinel import Sentinel
1414
from redis.exceptions import RedisError, WatchError
1515

16-
from .constants import default_queue_name, expires_extra_ms, job_key_prefix, result_key_prefix
16+
from .constants import (
17+
DEFAULT_PRIORITY,
18+
MAX_PRIORITY,
19+
MIN_PRIORITY,
20+
compute_ready_score,
21+
default_queue_name,
22+
deferred_queue_key,
23+
expires_extra_ms,
24+
job_key_prefix,
25+
ready_queue_key,
26+
result_key_prefix,
27+
)
1728
from .jobs import Deserializer, Job, JobDef, JobResult, Serializer, deserialize_job, serialize_job
1829
from .utils import timestamp_ms, to_ms, to_unix_ms
1930

@@ -126,6 +137,7 @@ async def enqueue_job(
126137
_defer_by: Union[None, int, float, timedelta] = None,
127138
_expires: Union[None, int, float, timedelta] = None,
128139
_job_try: Optional[int] = None,
140+
_priority: int = DEFAULT_PRIORITY,
129141
**kwargs: Any,
130142
) -> Optional[Job]:
131143
"""
@@ -140,9 +152,13 @@ async def enqueue_job(
140152
:param _expires: do not start or retry a job after this duration;
141153
defaults to 24 hours plus deferring time, if any
142154
:param _job_try: useful when re-enqueueing jobs within a job
155+
:param _priority: integer priority in the range 1..10 (1 = highest, 10 = lowest);
156+
defaults to 5. Within the same priority, jobs run FIFO by enqueue time.
143157
:param kwargs: any keyword arguments to pass to the function
144158
:return: :class:`arq.jobs.Job` instance or ``None`` if a job with this ID already exists
145159
"""
160+
if not (MIN_PRIORITY <= _priority <= MAX_PRIORITY):
161+
raise ValueError(f'_priority must be between {MIN_PRIORITY} and {MAX_PRIORITY}, got {_priority}')
146162
if _queue_name is None:
147163
_queue_name = self.default_queue_name
148164
job_id = _job_id or uuid4().hex
@@ -161,25 +177,80 @@ async def enqueue_job(
161177

162178
enqueue_time_ms = timestamp_ms()
163179
if _defer_until is not None:
164-
score = to_unix_ms(_defer_until)
180+
run_at_ms = to_unix_ms(_defer_until)
181+
deferred = True
165182
elif defer_by_ms:
166-
score = enqueue_time_ms + defer_by_ms
183+
run_at_ms = enqueue_time_ms + defer_by_ms
184+
deferred = True
167185
else:
168-
score = enqueue_time_ms
169-
170-
expires_ms = expires_ms or score - enqueue_time_ms + self.expires_extra_ms
171-
172-
job = serialize_job(function, args, kwargs, _job_try, enqueue_time_ms, serializer=self.job_serializer)
186+
run_at_ms = enqueue_time_ms
187+
deferred = False
188+
189+
expires_ms = expires_ms or run_at_ms - enqueue_time_ms + self.expires_extra_ms
190+
191+
job = serialize_job(
192+
function,
193+
args,
194+
kwargs,
195+
_job_try,
196+
enqueue_time_ms,
197+
priority=_priority,
198+
serializer=self.job_serializer,
199+
)
173200
pipe.multi()
174201
pipe.psetex(job_key, expires_ms, job)
175-
pipe.zadd(_queue_name, {job_id: score})
202+
if deferred:
203+
pipe.zadd(deferred_queue_key(_queue_name), {job_id: run_at_ms})
204+
else:
205+
pipe.zadd(
206+
ready_queue_key(_queue_name),
207+
{job_id: compute_ready_score(_priority, enqueue_time_ms)},
208+
)
176209
try:
177210
await pipe.execute()
178211
except WatchError:
179212
# job got enqueued since we checked 'job_exists'
180213
return None
181214
return Job(job_id, redis=self, _queue_name=_queue_name, _deserializer=self.job_deserializer)
182215

216+
async def queue_size(self, queue_name: Optional[str] = None) -> int:
217+
"""
218+
Total number of jobs in this queue across both the ready and deferred sub-zsets.
219+
"""
220+
if queue_name is None:
221+
queue_name = self.default_queue_name
222+
async with self.pipeline(transaction=False) as pipe:
223+
pipe.zcard(ready_queue_key(queue_name))
224+
pipe.zcard(deferred_queue_key(queue_name))
225+
ready, deferred = await pipe.execute()
226+
return int(ready) + int(deferred)
227+
228+
async def job_queue_score(
229+
self, queue_name: str, job_id: str
230+
) -> tuple[Optional[int], bool]:
231+
"""
232+
Look up a job's score in the queue. Returns ``(score, is_deferred)``;
233+
``score`` is ``None`` when the job is in neither sub-zset.
234+
"""
235+
async with self.pipeline(transaction=False) as pipe:
236+
pipe.zscore(ready_queue_key(queue_name), job_id)
237+
pipe.zscore(deferred_queue_key(queue_name), job_id)
238+
ready_score, deferred_score = await pipe.execute()
239+
if ready_score is not None:
240+
return int(ready_score), False
241+
if deferred_score is not None:
242+
return int(deferred_score), True
243+
return None, False
244+
245+
@staticmethod
246+
def zrem_from_queue(pipe: Any, queue_name: str, job_id: str) -> None:
247+
"""
248+
Stage ``ZREM`` against both sub-zsets on an existing pipeline.
249+
Idempotent — safe even if the job only lives in one of them.
250+
"""
251+
pipe.zrem(ready_queue_key(queue_name), job_id)
252+
pipe.zrem(deferred_queue_key(queue_name), job_id)
253+
183254
async def _get_job_result(self, key: bytes) -> JobResult:
184255
job_id = key[len(result_key_prefix) :].decode()
185256
job = Job(job_id, self, _deserializer=self.job_deserializer)
@@ -209,11 +280,16 @@ async def _get_job_def(self, job_id: bytes, score: int) -> JobDef:
209280

210281
async def queued_jobs(self, *, queue_name: Optional[str] = None) -> list[JobDef]:
211282
"""
212-
Get information about queued, mostly useful when testing.
283+
Get information about queued jobs across both the ready and deferred sub-zsets.
284+
Mostly useful when testing.
213285
"""
214286
if queue_name is None:
215287
queue_name = self.default_queue_name
216-
jobs = await self.zrange(queue_name, withscores=True, start=0, end=-1)
288+
async with self.pipeline(transaction=False) as pipe:
289+
pipe.zrange(ready_queue_key(queue_name), withscores=True, start=0, end=-1)
290+
pipe.zrange(deferred_queue_key(queue_name), withscores=True, start=0, end=-1)
291+
ready, deferred = await pipe.execute()
292+
jobs = list(ready) + list(deferred)
217293
return await asyncio.gather(*[self._get_job_def(job_id, int(score)) for job_id, score in jobs])
218294

219295

arq/constants.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,3 +16,24 @@
1616

1717
# extra time after the job is expected to start when the job key should expire, 1 day in ms
1818
expires_extra_ms = 86_400_000
19+
20+
# priority constants. ready-zset score = PRIORITY_FACTOR * priority + timestamp_ms,
21+
# so lower priority value runs first; FIFO is preserved within a tie via the timestamp.
22+
# 1e15 keeps the priority bucket strictly above any realistic timestamp_ms; doubles still
23+
# resolve ms-level differences at this magnitude (gap ~2.0 at ~1e16).
24+
PRIORITY_FACTOR = 10**15
25+
DEFAULT_PRIORITY = 5
26+
MIN_PRIORITY = 1
27+
MAX_PRIORITY = 10
28+
29+
30+
def ready_queue_key(queue_name: str) -> str:
31+
return f'{queue_name}:ready'
32+
33+
34+
def deferred_queue_key(queue_name: str) -> str:
35+
return f'{queue_name}:deferred'
36+
37+
38+
def compute_ready_score(priority: int, timestamp_ms: int) -> int:
39+
return PRIORITY_FACTOR * priority + timestamp_ms

arq/jobs.py

Lines changed: 77 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -2,14 +2,25 @@
22
import logging
33
import pickle
44
import warnings
5-
from dataclasses import dataclass
6-
from datetime import datetime
5+
from dataclasses import dataclass, field
6+
from datetime import datetime, timezone
77
from enum import Enum
88
from typing import Any, Callable, Optional
99

1010
from redis.asyncio import Redis
1111

12-
from .constants import abort_jobs_ss, default_queue_name, in_progress_key_prefix, job_key_prefix, result_key_prefix
12+
from .constants import (
13+
DEFAULT_PRIORITY,
14+
MIN_PRIORITY,
15+
abort_jobs_ss,
16+
compute_ready_score,
17+
default_queue_name,
18+
deferred_queue_key,
19+
in_progress_key_prefix,
20+
job_key_prefix,
21+
ready_queue_key,
22+
result_key_prefix,
23+
)
1324
from .utils import ms_to_datetime, poll, timestamp_ms
1425

1526
logger = logging.getLogger('arq.jobs')
@@ -39,6 +50,9 @@ class JobStatus(str, Enum):
3950
not_found = 'not_found'
4051

4152

53+
_EPOCH = datetime.fromtimestamp(0, tz=timezone.utc)
54+
55+
4256
@dataclass
4357
class JobDef:
4458
function: str
@@ -48,6 +62,9 @@ class JobDef:
4862
enqueue_time: datetime
4963
score: Optional[int]
5064
job_id: Optional[str]
65+
# priority defaults so existing constructor sites that predate the priority feature keep working;
66+
# dataclass inheritance forces every JobResult-only field below to also carry a default.
67+
priority: int = DEFAULT_PRIORITY
5168

5269
def __post_init__(self) -> None:
5370
if isinstance(self.score, float):
@@ -56,11 +73,11 @@ def __post_init__(self) -> None:
5673

5774
@dataclass
5875
class JobResult(JobDef):
59-
success: bool
60-
result: Any
61-
start_time: datetime
62-
finish_time: datetime
63-
queue_name: str
76+
success: bool = False
77+
result: Any = None
78+
start_time: datetime = field(default_factory=lambda: _EPOCH)
79+
finish_time: datetime = field(default_factory=lambda: _EPOCH)
80+
queue_name: str = ''
6481

6582

6683
class Job:
@@ -104,8 +121,9 @@ async def result(
104121
async for delay in poll(poll_delay):
105122
async with self._redis.pipeline(transaction=True) as tr:
106123
tr.get(result_key_prefix + self.job_id)
107-
tr.zscore(self._queue_name, self.job_id)
108-
v, s = await tr.execute()
124+
tr.zscore(ready_queue_key(self._queue_name), self.job_id)
125+
tr.zscore(deferred_queue_key(self._queue_name), self.job_id)
126+
v, ready_score, deferred_score = await tr.execute()
109127

110128
if v:
111129
info = deserialize_result(v, deserializer=self._deserializer)
@@ -115,7 +133,7 @@ async def result(
115133
raise info.result
116134
else:
117135
raise SerializationError(info.result)
118-
elif s is None:
136+
elif ready_score is None and deferred_score is None:
119137
raise ResultNotFound(
120138
'Not waiting for job result because the job is not in queue. '
121139
'Is the worker function configured to keep result?'
@@ -134,10 +152,25 @@ async def info(self) -> Optional[JobDef]:
134152
if v:
135153
info = deserialize_job(v, deserializer=self._deserializer)
136154
if info:
137-
s = await self._redis.zscore(self._queue_name, self.job_id)
138-
info.score = None if s is None else int(s)
155+
score, _ = await self._lookup_queue_score()
156+
info.score = score
139157
return info
140158

159+
async def _lookup_queue_score(self) -> tuple[Optional[int], bool]:
160+
"""
161+
Return ``(score, is_deferred)`` for this job. ``score`` is ``None`` when the job
162+
isn't in either sub-zset.
163+
"""
164+
async with self._redis.pipeline(transaction=False) as pipe:
165+
pipe.zscore(ready_queue_key(self._queue_name), self.job_id)
166+
pipe.zscore(deferred_queue_key(self._queue_name), self.job_id)
167+
ready_score, deferred_score = await pipe.execute()
168+
if ready_score is not None:
169+
return int(ready_score), False
170+
if deferred_score is not None:
171+
return int(deferred_score), True
172+
return None, False
173+
141174
async def result_info(self) -> Optional[JobResult]:
142175
"""
143176
Information about the job result if available, does not wait for the result. Does not raise an exception
@@ -156,15 +189,18 @@ async def status(self) -> JobStatus:
156189
async with self._redis.pipeline(transaction=True) as tr:
157190
tr.exists(result_key_prefix + self.job_id)
158191
tr.exists(in_progress_key_prefix + self.job_id)
159-
tr.zscore(self._queue_name, self.job_id)
160-
is_complete, is_in_progress, score = await tr.execute()
192+
tr.zscore(ready_queue_key(self._queue_name), self.job_id)
193+
tr.zscore(deferred_queue_key(self._queue_name), self.job_id)
194+
is_complete, is_in_progress, ready_score, deferred_score = await tr.execute()
161195

162196
if is_complete:
163197
return JobStatus.complete
164198
elif is_in_progress:
165199
return JobStatus.in_progress
166-
elif score:
167-
return JobStatus.deferred if score > timestamp_ms() else JobStatus.queued
200+
elif ready_score is not None:
201+
return JobStatus.queued
202+
elif deferred_score is not None:
203+
return JobStatus.deferred
168204
else:
169205
return JobStatus.not_found
170206

@@ -177,11 +213,16 @@ async def abort(self, *, timeout: Optional[float] = None, poll_delay: float = 0.
177213
:param poll_delay: how often to poll redis for the job result
178214
:return: True if the job aborted properly, False otherwise
179215
"""
180-
job_info = await self.info()
181-
if job_info and job_info.score and job_info.score > timestamp_ms():
216+
# if the job is currently sitting in the deferred sub-zset, hoist it into ready at
217+
# top priority so the worker picks it up on the next poll and observes the abort flag.
218+
_, is_deferred = await self._lookup_queue_score()
219+
if is_deferred:
182220
async with self._redis.pipeline(transaction=True) as tr:
183-
tr.zrem(self._queue_name, self.job_id)
184-
tr.zadd(self._queue_name, {self.job_id: 1})
221+
tr.zrem(deferred_queue_key(self._queue_name), self.job_id)
222+
tr.zadd(
223+
ready_queue_key(self._queue_name),
224+
{self.job_id: compute_ready_score(MIN_PRIORITY, timestamp_ms())},
225+
)
185226
await tr.execute()
186227

187228
await self._redis.zadd(abort_jobs_ss, {self.job_id: timestamp_ms()})
@@ -215,9 +256,17 @@ def serialize_job(
215256
job_try: Optional[int],
216257
enqueue_time_ms: int,
217258
*,
259+
priority: int = DEFAULT_PRIORITY,
218260
serializer: Optional[Serializer] = None,
219261
) -> bytes:
220-
data = {'t': job_try, 'f': function_name, 'a': args, 'k': kwargs, 'et': enqueue_time_ms}
262+
data = {
263+
't': job_try,
264+
'f': function_name,
265+
'a': args,
266+
'k': kwargs,
267+
'et': enqueue_time_ms,
268+
'priority': priority,
269+
}
221270
if serializer is None:
222271
serializer = pickle.dumps
223272
try:
@@ -240,6 +289,7 @@ def serialize_result(
240289
queue_name: str,
241290
job_id: str,
242291
*,
292+
priority: int = DEFAULT_PRIORITY,
243293
serializer: Optional[Serializer] = None,
244294
) -> Optional[bytes]:
245295
data = {
@@ -254,6 +304,7 @@ def serialize_result(
254304
'ft': finished_ms,
255305
'q': queue_name,
256306
'id': job_id,
307+
'priority': priority,
257308
}
258309
if serializer is None:
259310
serializer = pickle.dumps
@@ -284,19 +335,20 @@ def deserialize_job(r: bytes, *, deserializer: Optional[Deserializer] = None) ->
284335
enqueue_time=ms_to_datetime(d['et']),
285336
score=None,
286337
job_id=None,
338+
priority=d.get('priority', DEFAULT_PRIORITY),
287339
)
288340
except Exception as e:
289341
raise DeserializationError('unable to deserialize job') from e
290342

291343

292344
def deserialize_job_raw(
293345
r: bytes, *, deserializer: Optional[Deserializer] = None
294-
) -> tuple[str, tuple[Any, ...], dict[str, Any], int, int]:
346+
) -> tuple[str, tuple[Any, ...], dict[str, Any], int, int, int]:
295347
if deserializer is None:
296348
deserializer = pickle.loads
297349
try:
298350
d = deserializer(r)
299-
return d['f'], d['a'], d['k'], d['t'], d['et']
351+
return d['f'], d['a'], d['k'], d['t'], d['et'], d.get('priority', DEFAULT_PRIORITY)
300352
except Exception as e:
301353
raise DeserializationError('unable to deserialize job') from e
302354

@@ -319,6 +371,7 @@ def deserialize_result(r: bytes, *, deserializer: Optional[Deserializer] = None)
319371
finish_time=ms_to_datetime(d['ft']),
320372
queue_name=d.get('q', '<unknown>'),
321373
job_id=d.get('id', '<unknown>'),
374+
priority=d.get('priority', DEFAULT_PRIORITY),
322375
)
323376
except Exception as e:
324377
raise DeserializationError('unable to deserialize job result') from e

0 commit comments

Comments
 (0)