feat(tasks): Integrate django-tasks for task management and progress tracking
This commit is contained in:
+110
-7
@@ -2,7 +2,85 @@ from time import sleep
|
||||
from django.core.mail import send_mail
|
||||
from django.http import HttpResponse
|
||||
from django.shortcuts import get_object_or_404
|
||||
from celery import shared_task
|
||||
from enum import Enum
|
||||
|
||||
try:
|
||||
from django_tasks import task
|
||||
HAS_DJANGO_TASKS = True
|
||||
except ImportError:
|
||||
from celery import shared_task
|
||||
from celery.result import AsyncResult
|
||||
|
||||
HAS_DJANGO_TASKS = False
|
||||
|
||||
class _CompatTaskResultStatus(Enum):
|
||||
READY = "READY"
|
||||
RUNNING = "RUNNING"
|
||||
FAILED = "FAILED"
|
||||
SUCCESSFUL = "SUCCESSFUL"
|
||||
|
||||
class _CompatError:
|
||||
def __init__(self, traceback):
|
||||
self.traceback = traceback
|
||||
|
||||
class _CompatTaskResult:
|
||||
def __init__(self, async_result):
|
||||
self._async_result = async_result
|
||||
|
||||
@property
|
||||
def id(self):
|
||||
return self._async_result.id
|
||||
|
||||
@property
|
||||
def status(self):
|
||||
state = self._async_result.state
|
||||
if state in ("PENDING",):
|
||||
return _CompatTaskResultStatus.READY
|
||||
if state in ("STARTED", "RETRY"):
|
||||
return _CompatTaskResultStatus.RUNNING
|
||||
if state == "SUCCESS":
|
||||
return _CompatTaskResultStatus.SUCCESSFUL
|
||||
if state == "FAILURE":
|
||||
return _CompatTaskResultStatus.FAILED
|
||||
return _CompatTaskResultStatus.RUNNING
|
||||
|
||||
def refresh(self):
|
||||
return self
|
||||
|
||||
@property
|
||||
def return_value(self):
|
||||
return self._async_result.result
|
||||
|
||||
@property
|
||||
def errors(self):
|
||||
if self._async_result.state == "FAILURE" and self._async_result.traceback:
|
||||
return [_CompatError(self._async_result.traceback)]
|
||||
return []
|
||||
|
||||
def task(func=None, *, takes_context=False, **kwargs):
|
||||
def decorator(inner_func):
|
||||
if takes_context:
|
||||
@shared_task(bind=True)
|
||||
def wrapped(self, *args, **inner_kwargs):
|
||||
class _Ctx:
|
||||
class _TaskResultRef:
|
||||
id = self.request.id
|
||||
|
||||
task_result = _TaskResultRef()
|
||||
attempt = 1
|
||||
|
||||
return inner_func(_Ctx(), *args, **inner_kwargs)
|
||||
else:
|
||||
wrapped = shared_task(inner_func)
|
||||
|
||||
wrapped.enqueue = wrapped.delay
|
||||
wrapped.get_result = lambda result_id: _CompatTaskResult(AsyncResult(result_id))
|
||||
return wrapped
|
||||
|
||||
if func is not None:
|
||||
return decorator(func)
|
||||
return decorator
|
||||
|
||||
from atlas.models import Case, Series, SeriesImage
|
||||
from generic.models import CimarCase
|
||||
from rad.settings import REMOTE_URL, CIMAR_USERNAME, CIMAR_PASSWORD
|
||||
@@ -10,10 +88,11 @@ from helpers.cimar import CimarAPI, NotFoundError
|
||||
from pydicom.uid import generate_uid
|
||||
from django.contrib.auth.models import User
|
||||
from django.core.files.base import ContentFile
|
||||
from django.core.cache import cache
|
||||
import copy
|
||||
import io
|
||||
|
||||
@shared_task()
|
||||
@task
|
||||
def push_case_to_cimar_task(case_id):
|
||||
"""Sends an email when the feedback form has been submitted."""
|
||||
case = get_object_or_404(Case, pk=case_id)
|
||||
@@ -56,9 +135,9 @@ def push_case_to_cimar_task(case_id):
|
||||
return 10
|
||||
|
||||
|
||||
@shared_task(bind=True)
|
||||
@task(takes_context=True)
|
||||
def series_reconstruct_task(
|
||||
self,
|
||||
context,
|
||||
series_id,
|
||||
user_id,
|
||||
recon_planes,
|
||||
@@ -71,6 +150,13 @@ def series_reconstruct_task(
|
||||
from loguru import logger
|
||||
from atlas import views as atlas_views
|
||||
|
||||
progress_key = f"series_reconstruct_progress:{context.task_result.id}"
|
||||
cache.set(
|
||||
progress_key,
|
||||
{"current": 0, "total": 0, "message": "Preparing reconstruction..."},
|
||||
timeout=60 * 60,
|
||||
)
|
||||
|
||||
series = get_object_or_404(Series, pk=series_id)
|
||||
user = get_object_or_404(User, pk=user_id)
|
||||
|
||||
@@ -225,6 +311,12 @@ def series_reconstruct_task(
|
||||
raise ValueError("No valid reconstruction planes selected")
|
||||
|
||||
total_slices = sum(len(v["recon_slices"]) for v in plane_slices.values())
|
||||
cache.set(
|
||||
progress_key,
|
||||
{"current": 0, "total": total_slices, "message": "Reconstruction started"},
|
||||
timeout=60 * 60,
|
||||
)
|
||||
|
||||
processed = 0
|
||||
created_series = []
|
||||
|
||||
@@ -266,13 +358,14 @@ def series_reconstruct_task(
|
||||
recon_image.save()
|
||||
|
||||
processed += 1
|
||||
self.update_state(
|
||||
state="PROGRESS",
|
||||
meta={
|
||||
cache.set(
|
||||
progress_key,
|
||||
{
|
||||
"current": processed,
|
||||
"total": total_slices,
|
||||
"message": f"Generating {plane_norm} reconstruction ({processed}/{total_slices})",
|
||||
},
|
||||
timeout=60 * 60,
|
||||
)
|
||||
|
||||
created_series.append(
|
||||
@@ -289,6 +382,16 @@ def series_reconstruct_task(
|
||||
len(created_series),
|
||||
)
|
||||
|
||||
cache.set(
|
||||
progress_key,
|
||||
{
|
||||
"current": total_slices,
|
||||
"total": total_slices,
|
||||
"message": "Reconstruction complete",
|
||||
},
|
||||
timeout=10 * 60,
|
||||
)
|
||||
|
||||
return {
|
||||
"series_id": series.pk,
|
||||
"created_series": created_series,
|
||||
|
||||
Reference in New Issue
Block a user