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.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
@@ -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
@@ -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 *
|
||||||
|
|||||||
@@ -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,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
|
||||||
|
|||||||
Reference in New Issue
Block a user