forked from ACMClassOJ/TesutoHime
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdispatch.py
More file actions
130 lines (118 loc) · 5.86 KB
/
Copy pathdispatch.py
File metadata and controls
130 lines (118 loc) · 5.86 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
from asyncio import FIRST_COMPLETED, CancelledError, create_task, sleep, wait
from logging import getLogger
from time import time
from typing import Awaitable, Callable, Optional
from uuid import uuid4
from typing_extensions import overload
import commons.task_typing
from commons.task_typing import (CompileResult, CompileTask, Input,
JudgeResult, JudgeTask, StatusUpdate,
StatusUpdateDone, StatusUpdateError,
StatusUpdateProgress, StatusUpdateStarted)
from commons.util import deserialize, serialize
from scheduler2.config import (redis, redis_queues,
task_concurrency_per_account, task_retries,
task_retry_interval_secs, task_timeout_secs)
from scheduler2.metrics import observe_runner, tasks_by_state, tasks_retried_total
from scheduler2.monitor import wait_until_offline
from scheduler2.util import (RateLimiter, RunnerOfflineException, TaskInfo,
taskinfo_from_task_id)
logger = getLogger(__name__)
classes = commons.task_typing.__dict__
rate_limiter = RateLimiter(task_concurrency_per_account)
@overload
async def run_task(
taskinfo: TaskInfo[CompileTask],
onprogress: Optional[Callable[[StatusUpdate], Awaitable]] = None,
rate_limit_group: Optional[str] = None,
retries_left: int = task_retries,
) -> CompileResult: pass
@overload
async def run_task(
taskinfo: TaskInfo[JudgeTask[Input]],
onprogress: Optional[Callable[[StatusUpdate], Awaitable]] = None,
rate_limit_group: Optional[str] = None,
retries_left: int = task_retries,
) -> JudgeResult: pass
async def run_task(taskinfo, onprogress = None, rate_limit_group = None,
retries_left = task_retries):
task_id = str(uuid4())
taskinfo.id = task_id
task = taskinfo.task
taskinfo_from_task_id[task_id] = taskinfo
tasks_by_state.labels(state='waiting_for_rate_limit').inc()
reached_queued = False
reached_started = False
async def retry(msg):
if retries_left <= 0:
logger.warn('task %(id)s failed: %(message)s', { 'id': task_id, 'message': msg }, 'task:fail')
tasks_retried_total.labels(reason='exhausted').inc()
raise Exception(msg)
logger.info('task %(id)s failed: %(message)s, retrying', { 'id': task_id, 'message': msg }, 'task:retry')
tasks_retried_total.labels(reason='retry').inc()
await sleep(task_retry_interval_secs)
return await run_task(taskinfo, onprogress, rate_limit_group,
retries_left - 1)
try:
async with rate_limiter.limit(rate_limit_group):
tasks_by_state.labels(state='waiting_for_rate_limit').dec()
tasks_by_state.labels(state='queued').inc()
reached_queued = True
logger.debug('running task %(id)s: %(task)s', { 'id': task_id, 'task': task }, 'task:start')
queues = redis_queues.task(task_id)
await redis.lpush(queues.task, serialize(task))
await redis.expire(queues.task, task_timeout_secs)
await redis.lpush(redis_queues.tasks_group(taskinfo.group), task_id)
task_timeout = time() + task_timeout_secs
offline_task = None
try:
while True:
progress_task = create_task(redis.brpop(queues.progress,
int(task_timeout - time())))
tasks = (progress_task,)
if offline_task is not None:
tasks = (progress_task, offline_task)
done, _ = await wait(tasks, return_when=FIRST_COMPLETED)
for done_task in done:
try:
res = await done_task
except RunnerOfflineException:
return await retry('Runner offline')
if res is None:
return await retry('Task timed out')
_, status = res
status: StatusUpdate = deserialize(status) # type: ignore
logger.debug('received status update from task %(id)s: %(status)s', { 'id': task_id, 'status': status }, 'task:update')
if isinstance(status, StatusUpdateStarted):
tasks_by_state.labels(state='queued').dec()
tasks_by_state.labels(state='started').inc()
reached_started = True
observe_runner(status.id)
offline_task = wait_until_offline(status.id)
if onprogress is not None:
await onprogress(status)
elif isinstance(status, StatusUpdateProgress):
if onprogress is not None:
await onprogress(status)
elif isinstance(status, StatusUpdateDone):
return status.result
elif isinstance(status, StatusUpdateError):
message = status.message
msg = f'Runner error: {message}'
logger.error(msg)
return await retry(msg)
else:
raise Exception(f'Unknown message from runner: {status}')
except CancelledError:
logger.info('aborting task %(id)s', { 'id': task_id }, 'task:abort')
await redis.lpush(queues.abort, 1)
await redis.expire(queues.abort, task_timeout_secs)
raise
finally:
if reached_started:
tasks_by_state.labels(state='started').dec()
elif reached_queued:
tasks_by_state.labels(state='queued').dec()
else:
tasks_by_state.labels(state='waiting_for_rate_limit').dec()
del taskinfo_from_task_id[task_id]