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.core.mail import send_mail
from django.http import HttpResponse from django.http import HttpResponse
from django.shortcuts import get_object_or_404 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 atlas.models import Case, Series, SeriesImage
from generic.models import CimarCase from generic.models import CimarCase
from rad.settings import REMOTE_URL, CIMAR_USERNAME, CIMAR_PASSWORD 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 pydicom.uid import generate_uid
from django.contrib.auth.models import User from django.contrib.auth.models import User
from django.core.files.base import ContentFile from django.core.files.base import ContentFile
from django.core.cache import cache
import copy import copy
import io import io
@shared_task() @task
def push_case_to_cimar_task(case_id): def push_case_to_cimar_task(case_id):
"""Sends an email when the feedback form has been submitted.""" """Sends an email when the feedback form has been submitted."""
case = get_object_or_404(Case, pk=case_id) case = get_object_or_404(Case, pk=case_id)
@@ -56,9 +135,9 @@ def push_case_to_cimar_task(case_id):
return 10 return 10
@shared_task(bind=True) @task(takes_context=True)
def series_reconstruct_task( def series_reconstruct_task(
self, context,
series_id, series_id,
user_id, user_id,
recon_planes, recon_planes,
@@ -71,6 +150,13 @@ def series_reconstruct_task(
from loguru import logger from loguru import logger
from atlas import views as atlas_views 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) series = get_object_or_404(Series, pk=series_id)
user = get_object_or_404(User, pk=user_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") raise ValueError("No valid reconstruction planes selected")
total_slices = sum(len(v["recon_slices"]) for v in plane_slices.values()) 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 processed = 0
created_series = [] created_series = []
@@ -266,13 +358,14 @@ def series_reconstruct_task(
recon_image.save() recon_image.save()
processed += 1 processed += 1
self.update_state( cache.set(
state="PROGRESS", progress_key,
meta={ {
"current": processed, "current": processed,
"total": total_slices, "total": total_slices,
"message": f"Generating {plane_norm} reconstruction ({processed}/{total_slices})", "message": f"Generating {plane_norm} reconstruction ({processed}/{total_slices})",
}, },
timeout=60 * 60,
) )
created_series.append( created_series.append(
@@ -289,6 +382,16 @@ def series_reconstruct_task(
len(created_series), len(created_series),
) )
cache.set(
progress_key,
{
"current": total_slices,
"total": total_slices,
"message": "Reconstruction complete",
},
timeout=10 * 60,
)
return { return {
"series_id": series.pk, "series_id": series.pk,
"created_series": created_series, "created_series": created_series,
+39 -28
View File
@@ -164,7 +164,16 @@ from .filters import (
) )
from .tasks import push_case_to_cimar_task, series_reconstruct_task from .tasks import push_case_to_cimar_task, series_reconstruct_task
from celery.result import AsyncResult try:
from django_tasks import TaskResultStatus
except ImportError:
from enum import Enum
class TaskResultStatus(Enum):
READY = "READY"
RUNNING = "RUNNING"
FAILED = "FAILED"
SUCCESSFUL = "SUCCESSFUL"
from django_tables2 import SingleTableView, SingleTableMixin from django_tables2 import SingleTableView, SingleTableMixin
from django_filters.views import FilterView from django_filters.views import FilterView
@@ -685,7 +694,7 @@ def series_optimize_htmx(request, series_id):
return HttpResponse('<div class="alert alert-danger mb-0">Slice thickness must be greater than 0.</div>') return HttpResponse('<div class="alert alert-danger mb-0">Slice thickness must be greater than 0.</div>')
if slice_spacing_val is not None and slice_spacing_val <= 0: if slice_spacing_val is not None and slice_spacing_val <= 0:
return HttpResponse('<div class="alert alert-danger mb-0">Slice spacing must be greater than 0.</div>') return HttpResponse('<div class="alert alert-danger mb-0">Slice spacing must be greater than 0.</div>')
async_task = series_reconstruct_task.delay( task_result = series_reconstruct_task.enqueue(
series_id=series.pk, series_id=series.pk,
user_id=request.user.pk, user_id=request.user.pk,
recon_planes=recon_planes, recon_planes=recon_planes,
@@ -699,7 +708,7 @@ def series_optimize_htmx(request, series_id):
'<div class="alert alert-info mb-2">' '<div class="alert alert-info mb-2">'
'Reconstruction queued. This runs in the background to avoid request timeout.' 'Reconstruction queued. This runs in the background to avoid request timeout.'
'</div>' '</div>'
f'<div id="recon-task-status" hx-get="{reverse("atlas:series_reconstruct_status", kwargs={"series_id": series.pk, "task_id": async_task.id})}" ' f'<div id="recon-task-status" hx-get="{reverse("atlas:series_reconstruct_status", kwargs={"series_id": series.pk, "task_id": task_result.id})}" '
'hx-trigger="load, every 2s" hx-swap="outerHTML">' 'hx-trigger="load, every 2s" hx-swap="outerHTML">'
'<div class="d-flex align-items-center gap-2"><span class="spinner-border spinner-border-sm text-primary" role="status"></span>' '<div class="d-flex align-items-center gap-2"><span class="spinner-border spinner-border-sm text-primary" role="status"></span>'
'<span class="small text-muted">Starting reconstruction task...</span></div></div>' '<span class="small text-muted">Starting reconstruction task...</span></div></div>'
@@ -717,29 +726,26 @@ def series_reconstruct_status_htmx(request, series_id, task_id):
if not series.check_user_can_edit(request.user): if not series.check_user_can_edit(request.user):
return HttpResponse('<div class="alert alert-danger mb-0">Permission denied</div>') return HttpResponse('<div class="alert alert-danger mb-0">Permission denied</div>')
task_result = AsyncResult(task_id) try:
state = task_result.state task_result = series_reconstruct_task.get_result(task_id)
except NotImplementedError:
return HttpResponse(
'<div id="recon-task-status" class="alert alert-warning mb-0">Task backend does not support result polling. Configure a backend with result retrieval support.</div>'
)
except Exception:
return HttpResponse('<div id="recon-task-status" class="alert alert-danger mb-0">Task result not found.</div>')
task_result.refresh()
state = task_result.status
poll_url = reverse("atlas:series_reconstruct_status", kwargs={"series_id": series.pk, "task_id": task_id}) poll_url = reverse("atlas:series_reconstruct_status", kwargs={"series_id": series.pk, "task_id": task_id})
progress = cache.get(f"series_reconstruct_progress:{task_id}") or {}
if state in ("PENDING", "STARTED", "RETRY"): current = int(progress.get("current", 0) or 0)
return HttpResponse( total = int(progress.get("total", 0) or 0)
( message = escape(str(progress.get("message", "Reconstruction is running...")))
f'<div id="recon-task-status" hx-get="{poll_url}" hx-trigger="every 2s" hx-swap="outerHTML">'
'<div class="d-flex align-items-center gap-2">'
'<span class="spinner-border spinner-border-sm text-primary" role="status"></span>'
'<span class="small text-muted">Reconstruction is running...</span>'
'</div></div>'
)
)
if state == "PROGRESS":
meta = task_result.info or {}
current = int(meta.get("current", 0) or 0)
total = int(meta.get("total", 0) or 0)
message = escape(str(meta.get("message", "Reconstruction running...")))
pct = int((current / total) * 100) if total > 0 else 0 pct = int((current / total) * 100) if total > 0 else 0
if state in (TaskResultStatus.READY, TaskResultStatus.RUNNING):
return HttpResponse( return HttpResponse(
( (
f'<div id="recon-task-status" hx-get="{poll_url}" hx-trigger="every 2s" hx-swap="outerHTML">' f'<div id="recon-task-status" hx-get="{poll_url}" hx-trigger="every 2s" hx-swap="outerHTML">'
@@ -753,8 +759,11 @@ def series_reconstruct_status_htmx(request, series_id, task_id):
) )
) )
if state == "SUCCESS": if state == TaskResultStatus.SUCCESSFUL:
payload = task_result.result if isinstance(task_result.result, dict) else {} try:
payload = task_result.return_value if isinstance(task_result.return_value, dict) else {}
except ValueError:
payload = {}
created_series = payload.get("created_series", []) created_series = payload.get("created_series", [])
if created_series: if created_series:
@@ -774,14 +783,16 @@ def series_reconstruct_status_htmx(request, series_id, task_id):
return HttpResponse('<div id="recon-task-status" class="alert alert-warning mb-0">Reconstruction finished but no output series were created.</div>') return HttpResponse('<div id="recon-task-status" class="alert alert-warning mb-0">Reconstruction finished but no output series were created.</div>')
if state == "FAILURE": if state == TaskResultStatus.FAILED:
err = escape(str(task_result.result)) err = "Unknown failure"
if task_result.errors:
err = escape(task_result.errors[-1].traceback.splitlines()[-1])
return HttpResponse( return HttpResponse(
f'<div id="recon-task-status" class="alert alert-danger mb-0">Reconstruction failed: {err}</div>' f'<div id="recon-task-status" class="alert alert-danger mb-0">Reconstruction failed: {err}</div>'
) )
return HttpResponse( return HttpResponse(
f'<div id="recon-task-status" class="alert alert-secondary mb-0">Task state: {escape(state)}</div>' f'<div id="recon-task-status" class="alert alert-secondary mb-0">Task state: {escape(str(state))}</div>'
) )
@@ -8758,7 +8769,7 @@ def collection_reset_answers(request, exam_id: int):
def push_case_to_cimar(request, case_id): def push_case_to_cimar(request, case_id):
push_case_to_cimar_task.delay(case_id=case_id) push_case_to_cimar_task.enqueue(case_id=case_id)
return HttpResponse("Push started") return HttpResponse("Push started")
case = get_object_or_404(Case, pk=case_id) case = get_object_or_404(Case, pk=case_id)
+12 -2
View File
@@ -86,8 +86,7 @@ INSTALLED_APPS = [
'django_jsonforms', 'django_jsonforms',
'django_svelte_jsoneditor', 'django_svelte_jsoneditor',
'django_psutil_dash', 'django_psutil_dash',
'django_tasks',
] ]
MIDDLEWARE = [ MIDDLEWARE = [
@@ -382,6 +381,17 @@ CIMAR_PASSWORD = ""
CELERY_BROKER_URL = "redis://redis:6379" CELERY_BROKER_URL = "redis://redis:6379"
CELERY_RESULT_BACKEND = "redis://redis:6379" CELERY_RESULT_BACKEND = "redis://redis:6379"
# Django 6 task framework settings.
# Use DJANGO_TASK_BACKEND to select your backend implementation.
TASKS = {
"default": {
"BACKEND": os.environ.get(
"DJANGO_TASK_BACKEND",
"django_tasks.backends.immediate.ImmediateBackend",
),
}
}
try: try:
from .settings_local import * from .settings_local import *
+6
View File
@@ -24,3 +24,9 @@ EMAIL_BACKEND = "django.core.mail.backends.locmem.EmailBackend"
CACHES["default"] = { CACHES["default"] = {
"BACKEND": "django.core.cache.backends.locmem.LocMemCache", "BACKEND": "django.core.cache.backends.locmem.LocMemCache",
} }
TASKS = {
"default": {
"BACKEND": "django_tasks.backends.immediate.ImmediateBackend",
}
}
+1
View File
@@ -1,5 +1,6 @@
#Django==3.2.13 #Django==3.2.13
Django==6.0.1 Django==6.0.1
django-tasks
django_debug_toolbar django_debug_toolbar
django_jquery django_jquery
django_reversion django_reversion