|
14 | 14 | from .constants import TaskPriority, TaskState |
15 | 15 | from .task import Task |
16 | 16 | from .task_run import TaskRun |
17 | | -from .utils import WaitFor |
18 | 17 |
|
19 | 18 | ConsumerCallback = Callable[[TaskRun, "TaskManager"], None] |
20 | 19 | AsyncExecutor = Callable[..., Coroutine[Any, Any, None]] |
@@ -164,14 +163,9 @@ def num_concurrent_tasks_for(self, task_name: str) -> int: |
164 | 163 | """The number of concurrent tasks for a given task_name""" |
165 | 164 | return len(self._concurrent_tasks[task_name]) |
166 | 165 |
|
167 | | - async def queue_and_wait(self, task: str, **params: Any) -> TaskRun: |
168 | | - run_id = self.execute(task, **params).id |
169 | | - waitfor = WaitFor(run_id=run_id) |
170 | | - self.register_handler(f"end.{waitfor.run_id}", waitfor) |
171 | | - try: |
172 | | - return await waitfor.waiter |
173 | | - finally: |
174 | | - self.unregister_handler(f"end.{waitfor.run_id}") |
| 166 | + async def queue_and_wait(self, task: str, **params: Any) -> Any: |
| 167 | + """Execute a task by-passing the broker task queue and wait for result""" |
| 168 | + return await self.execute(task, **params).waiter |
175 | 169 |
|
176 | 170 | def execute(self, task: str, **params: Any) -> TaskRun: |
177 | 171 | """Execute a Task by-passing the broker task queue""" |
@@ -210,20 +204,17 @@ async def _consume_tasks(self) -> None: |
210 | 204 | task_run.start = microseconds() |
211 | 205 | task_run.set_state(TaskState.running) |
212 | 206 | task_context = task_run.task.create_context(self, task_run=task_run) |
213 | | - info = await self.broker.get_tasks_info(task_name) |
214 | | - if not info[0].enabled: |
215 | | - task_run.set_state(TaskState.aborted) |
216 | | - task_run.waiter.set_result(None) |
| 207 | + self._concurrent_tasks[task_name][task_run.id] = task_run |
217 | 208 | # |
218 | | - elif task_run.task.max_concurrency <= self.num_concurrent_tasks_for( |
219 | | - task_name |
220 | | - ): |
| 209 | + if task_run.task.max_concurrency < self.num_concurrent_tasks_for(task_name): |
221 | 210 | task_run.set_state(TaskState.rate_limited) |
222 | 211 | task_run.waiter.set_result(None) |
| 212 | + elif not (await self.broker.get_tasks_info(task_name))[0].enabled: |
| 213 | + task_run.set_state(TaskState.aborted) |
| 214 | + task_run.waiter.set_result(None) |
223 | 215 | # |
224 | 216 | else: |
225 | 217 | task_context.logger.info("start") |
226 | | - self._concurrent_tasks[task_name][task_run.id] = task_run |
227 | 218 | self.dispatch(task_run, "start") |
228 | 219 | try: |
229 | 220 | result = await task_run.task.executor(task_context) |
|
0 commit comments