import logging from django.utils import timezone from django.utils.module_loading import import_string from infrasynth.shared.settings_utils import get_setting from .models import ScheduledTask, TaskExecution from .signals import task_failed, task_scheduled, task_started logger = logging.getLogger(__name__) class TaskService: """Public API for scheduled/on-demand Celery task management.""" def run_now(self, task_id: int) -> TaskExecution: """Triggers an immediate execution of a scheduled task via Celery.""" task = ScheduledTask.objects.get(pk=task_id) execution = TaskExecution.objects.create( task=task, status=TaskExecution.Status.PENDING, started_at=timezone.now(), ) self._prune_history(task) try: task_func = import_string(task.task_path) except ImportError: logger.error("Could not import task path '%s'", task.task_path) execution.status = TaskExecution.Status.FAILURE execution.error_traceback = f"Could not import task path '{task.task_path}'" execution.completed_at = timezone.now() execution.save() task_failed.send( sender=TaskExecution, task_name=task.name, task_id=execution.id, error=execution.error_traceback, traceback="", ) return execution args = task.args or [] kwargs = dict(task.kwargs or {}) queue = task.queue or get_setting("INFRASYNTH_SCHEDULER", "DEFAULT_QUEUE", "default") from infrasynth.tenancy.context import tenant_context with tenant_context(task.tenant): return self._invoke(task, task_func, execution, args, kwargs, queue) def _invoke(self, task, task_func, execution, args, kwargs, queue): import inspect def accepts_tenant_id(func) -> bool: try: params = inspect.signature(func).parameters except (TypeError, ValueError): return False return "tenant_id" in params or any(p.kind is p.VAR_KEYWORD for p in params.values()) if accepts_tenant_id(task_func): kwargs.setdefault("tenant_id", str(task.tenant_id)) try: apply_async = getattr(task_func, "apply_async", None) if callable(apply_async): async_result = task_func.apply_async( args=args, kwargs=kwargs, queue=queue, priority=task.priority, ) execution.celery_task_id = str(async_result.id) execution.status = TaskExecution.Status.RUNNING execution.save(update_fields=["celery_task_id", "status"]) else: execution.celery_task_id = "" if args or kwargs: task_func(*args, **kwargs) else: task_func() execution.status = TaskExecution.Status.SUCCESS execution.completed_at = timezone.now() execution.save(update_fields=["celery_task_id", "status", "completed_at"]) except Exception as exc: # noqa: BLE001 logger.exception("Failed to trigger task '%s'", task.name) execution.status = TaskExecution.Status.FAILURE execution.error_traceback = str(exc) execution.completed_at = timezone.now() execution.save() task_failed.send( sender=TaskExecution, tenant_id=str(task.tenant_id), task_name=task.name, task_id=execution.id, error=str(exc), traceback="", ) return execution task_scheduled.send(sender=TaskExecution, tenant_id=str(task.tenant_id), task_name=task.name, eta=None) task_started.send( sender=TaskExecution, tenant_id=str(task.tenant_id), task_name=task.name, task_id=execution.id, worker=None, ) return execution def toggle(self, task_id: int) -> ScheduledTask: """Enables or disables a scheduled task.""" task = ScheduledTask.objects.get(pk=task_id) task.is_active = not task.is_active task.save(update_fields=["is_active"]) return task @staticmethod def _prune_history(task: ScheduledTask) -> None: """Keeps the newest ``MAX_EXECUTION_HISTORY_PER_TASK`` executions per task.""" limit = int(get_setting("INFRASYNTH_SCHEDULER", "MAX_EXECUTION_HISTORY_PER_TASK", 1000)) if limit <= 0: return keep = list( TaskExecution.objects.filter(task=task).order_by("-started_at", "-id").values_list("id", flat=True)[:limit] ) if keep: TaskExecution.objects.filter(task=task).exclude(id__in=keep).delete() def get_queue_status(self) -> dict: """Returns active/scheduled/reserved task counts per queue.""" from celery import current_app inspect = current_app.control.inspect() stats: dict = {"queues": {}, "total_active": 0, "total_scheduled": 0, "total_reserved": 0} active = inspect.active() or {} scheduled = inspect.scheduled() or {} reserved = inspect.reserved() or {} for worker, tasks in active.items(): for t in tasks: queue = t.get("delivery_info", {}).get("routing_key", "default") entry = stats["queues"].setdefault(queue, {"active": 0, "scheduled": 0, "reserved": 0}) entry["active"] += 1 stats["total_active"] += 1 for worker, tasks in scheduled.items(): for t in tasks: queue = t.get("delivery_info", {}).get("routing_key", "default") entry = stats["queues"].setdefault(queue, {"active": 0, "scheduled": 0, "reserved": 0}) entry["scheduled"] += 1 stats["total_scheduled"] += 1 for worker, tasks in reserved.items(): for t in tasks: queue = t.get("delivery_info", {}).get("routing_key", "default") entry = stats["queues"].setdefault(queue, {"active": 0, "scheduled": 0, "reserved": 0}) entry["reserved"] += 1 stats["total_reserved"] += 1 return stats def get_workers(self) -> list[dict]: """Returns active workers and their stats.""" from celery import current_app inspect = current_app.control.inspect() workers = [] stats = inspect.stats() or {} active = inspect.active() or {} for hostname, info in stats.items(): workers.append( { "hostname": hostname, "active_tasks": len(active.get(hostname, []) or []), "processed": info.get("total", {}).get("task", 0), "uptime_seconds": info.get("uptime", 0), "queues": info.get("queues", []), "status": "online", } ) return workers