feat(tasks): Integrate django-tasks for task management and progress tracking

This commit is contained in:
Ross
2026-05-17 12:41:17 +01:00
parent 7ed0a8f378
commit e6a20bb851
5 changed files with 169 additions and 38 deletions
+110 -7
View File
@@ -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,