diff --git a/api/changelog.d/celery-task-kwargs-parsing.fixed.md b/api/changelog.d/celery-task-kwargs-parsing.fixed.md new file mode 100644 index 0000000000..25403ee7b9 --- /dev/null +++ b/api/changelog.d/celery-task-kwargs-parsing.fixed.md @@ -0,0 +1 @@ +`task_args` serialization no longer returns HTTP 500 errors when Celery truncates stored task keyword arguments diff --git a/api/src/backend/api/celery_utils.py b/api/src/backend/api/celery_utils.py new file mode 100644 index 0000000000..6e8d825d70 --- /dev/null +++ b/api/src/backend/api/celery_utils.py @@ -0,0 +1,40 @@ +import ast +import json +from typing import Any + +_UNPARSED = object() + + +def decode_celery_field(value: Any, default: Any) -> Any: + """Decode a Celery result field and require JSON-serializable output.""" + decoded = value + for _ in range(2): + if not isinstance(decoded, str): + break + + text = decoded.strip() + if not text: + decoded = default + break + + parsed = _UNPARSED + for parser in (json.loads, ast.literal_eval): + try: + parsed = parser(text) + break + except (TypeError, ValueError, SyntaxError): + continue + + if parsed is _UNPARSED: + raise ValueError("Unable to decode Celery result field") + decoded = parsed + + decoded = default if decoded is None else decoded + try: + json.dumps(decoded, allow_nan=False) + except (TypeError, ValueError) as error: + raise ValueError( + "Decoded Celery result field is not JSON serializable" + ) from error + + return decoded diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index 7c2a03b965..3dd96d162b 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -65,6 +65,7 @@ from api.v1.views import ( ) from botocore.exceptions import ClientError, NoCredentialsError from celery import states +from celery.utils.saferepr import saferepr from conftest import ( API_JSON_CONTENT_TYPE, TEST_PASSWORD, @@ -5043,11 +5044,60 @@ class TestTaskViewSet: reverse("task-detail", kwargs={"pk": task1.id}), ) assert response.status_code == status.HTTP_200_OK + assert response.json()["data"]["attributes"]["task_args"] == { + "kwarg1": "value1" + } assert ( response.json()["data"]["attributes"]["name"] == task1.task_runner_task.task_name ) + def test_tasks_retrieve_hides_tenant_id( + self, authenticated_client, tasks_fixture, tenants_fixture + ): + task, *_ = tasks_fixture + task.task_runner_task.task_kwargs = json.dumps( + repr( + { + "tenant_id": str(tenants_fixture[0].id), + "enabled": True, + "scan_id": None, + "label": "True North", + } + ) + ) + task.task_runner_task.save(update_fields=["task_kwargs"]) + + response = authenticated_client.get( + reverse("task-detail", kwargs={"pk": task.id}), + ) + + assert response.status_code == status.HTTP_200_OK + assert response.json()["data"]["attributes"]["task_args"] == { + "enabled": True, + "scan_id": None, + "label": "True North", + } + + def test_tasks_retrieve_with_truncated_kwargs_returns_empty_task_args( + self, authenticated_client, tasks_fixture + ): + task, *_ = tasks_fixture + kwargs_repr = saferepr( + {"finding_ids": [str(uuid4()) for _ in range(30)]}, maxlen=1024 + ) + assert "..." in kwargs_repr + task.task_runner_task.task_kwargs = json.dumps(kwargs_repr) + task.task_runner_task.save(update_fields=["task_kwargs"]) + + response = authenticated_client.get( + reverse("task-detail", kwargs={"pk": task.id}), + ) + + assert response.status_code == status.HTTP_200_OK + assert response.headers["Content-Type"] == API_JSON_CONTENT_TYPE + assert response.json()["data"]["attributes"]["task_args"] == {} + def test_tasks_invalid_retrieve(self, authenticated_client): response = authenticated_client.get( reverse("task-detail", kwargs={"pk": "invalid_id"}) diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 9292a0c633..b6f600098c 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -1,8 +1,10 @@ import base64 import json +import logging from datetime import UTC, datetime, timedelta import yaml +from api.celery_utils import decode_celery_field from api.db_router import MainRouter from api.exceptions import ConflictException from api.models import ( @@ -59,6 +61,7 @@ from api.v1.serializer_utils.lighthouse import ( from api.v1.serializer_utils.processors import ProcessorConfigField from api.v1.serializer_utils.providers import ProviderSecretField from api.validators import validate_lighthouse_openai_compatible_base_url +from config.custom_logging import BackendLogger from django.conf import settings from django.contrib.auth import authenticate from django.contrib.auth.models import update_last_login @@ -79,6 +82,8 @@ from rest_framework_simplejwt.settings import api_settings from rest_framework_simplejwt.tokens import RefreshToken from rest_framework_simplejwt.utils import get_md5_hash_password +logger = logging.getLogger(BackendLogger.API) + # Base @@ -607,13 +612,24 @@ class TaskSerializer(RLSSerializer, TaskBase): @extend_schema_field(serializers.JSONField()) def get_task_args(self, obj): - task_args = self.get_json_field(obj, "task_kwargs") - # Celery task_kwargs are stored as a double string JSON in the database when not empty - if isinstance(task_args, str): - task_args = json.loads(task_args.replace("'", '"').replace("None", "null")) - # Remove tenant_id from task_kwargs if present - task_args.pop("tenant_id", None) + task_kwargs = ( + getattr(obj.task_runner_task, "task_kwargs", None) + if obj.task_runner_task + else None + ) + try: + task_args = decode_celery_field(task_kwargs, {}) + if not isinstance(task_args, dict): + raise ValueError("Decoded task kwargs must be a dictionary") + except ValueError: + logger.warning( + "Unable to decode task kwargs for task %s; returning empty task_args.", + obj.id, + ) + return {} + task_args = task_args.copy() + task_args.pop("tenant_id", None) return task_args @staticmethod diff --git a/api/src/backend/tasks/jobs/orphan_recovery.py b/api/src/backend/tasks/jobs/orphan_recovery.py index 05a2083dbe..c8cda54cd2 100644 --- a/api/src/backend/tasks/jobs/orphan_recovery.py +++ b/api/src/backend/tasks/jobs/orphan_recovery.py @@ -18,12 +18,11 @@ This is the shared engine behind both the periodic Beat watchdog and the `reconcile_orphan_tasks` management command. """ -import ast -import json from contextlib import contextmanager from datetime import UTC, datetime, timedelta from uuid import uuid4 +from api.celery_utils import decode_celery_field from celery import current_app, states from celery.utils.log import get_task_logger from django.db import connections @@ -138,34 +137,6 @@ def revoke_task(task_result, terminate: bool = True) -> None: logger.exception(f"Failed to revoke task {task_result.task_id}") -def _decode_celery_field(value, default): - """Decode django-celery-results' stored task_args/task_kwargs to a Python object. - - The backend stores them as a (sometimes double-encoded) repr/JSON string. An - empty or missing field returns ``default``; a non-empty value that cannot be - decoded raises ``ValueError`` so the caller can avoid re-enqueuing a task with - the wrong arguments. - """ - obj = value - for _ in range(2): # values can be double-encoded (a string holding a repr) - if not isinstance(obj, str): - break - text = obj.strip() - if not text: - return default - parsed = None - for parser in (ast.literal_eval, json.loads): - try: - parsed = parser(text) - break - except (ValueError, SyntaxError, TypeError): - continue - if parsed is None: - raise ValueError(f"undecodable celery field: {text[:120]!r}") - obj = parsed - return default if obj is None else obj - - def reconcile_orphans( grace_minutes: int = 2, max_attempts: int = 3, @@ -313,8 +284,10 @@ def _recover_task(task_result, max_attempts: int, window_hours: int) -> str: return "failed" try: - args = _decode_celery_field(args_repr, []) - kwargs = _decode_celery_field(kwargs_repr, {}) + args = decode_celery_field(args_repr, []) + kwargs = decode_celery_field(kwargs_repr, {}) + if not isinstance(args, (list, tuple)) or not isinstance(kwargs, dict): + raise ValueError("Stored task arguments have invalid types") except ValueError: logger.error( "Orphan %s (%s): could not decode stored args/kwargs, not re-enqueuing", @@ -324,8 +297,8 @@ def _recover_task(task_result, max_attempts: int, window_hours: int) -> str: return "failed" new_task_id = str(uuid4()) task_obj.apply_async( - args=list(args) if isinstance(args, (list, tuple)) else [], - kwargs=kwargs if isinstance(kwargs, dict) else {}, + args=list(args), + kwargs=kwargs, task_id=new_task_id, ) logger.info( diff --git a/api/src/backend/tasks/tests/test_orphan_recovery.py b/api/src/backend/tasks/tests/test_orphan_recovery.py index 074e311b70..77ed831f18 100644 --- a/api/src/backend/tasks/tests/test_orphan_recovery.py +++ b/api/src/backend/tasks/tests/test_orphan_recovery.py @@ -3,12 +3,13 @@ from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest +from api.celery_utils import decode_celery_field from celery import states +from celery.utils.saferepr import saferepr from django.test import override_settings from django_celery_results.models import TaskResult from tasks.jobs.orphan_recovery import ( _SKIP_RECOVERY, - _decode_celery_field, _reconcile_task_results, _recovery_attempt_count, advisory_lock, @@ -36,24 +37,77 @@ def _orphan_result(*, name, kwargs, worker, created_minutes_ago, status=states.S return tr -@pytest.mark.django_db class TestDecodeCeleryField: + def test_decodes_strict_json(self): + assert decode_celery_field('{"enabled": true, "scan_id": null}', {}) == { + "enabled": True, + "scan_id": None, + } + def test_decodes_single_encoded_repr(self): - assert _decode_celery_field("{'tenant_id': 'abc'}", {}) == {"tenant_id": "abc"} + assert decode_celery_field("{'tenant_id': 'abc'}", {}) == {"tenant_id": "abc"} def test_decodes_double_encoded(self): import json stored = json.dumps(repr({"tenant_id": "abc", "scan_id": "s1"})) - assert _decode_celery_field(stored, {}) == {"tenant_id": "abc", "scan_id": "s1"} + assert decode_celery_field(stored, {}) == { + "tenant_id": "abc", + "scan_id": "s1", + } + + def test_python_words_inside_strings_are_preserved(self): + stored = repr( + { + "enabled": True, + "scan_id": None, + "label": "True North", + "note": "None", + } + ) + + assert decode_celery_field(stored, {}) == { + "enabled": True, + "scan_id": None, + "label": "True North", + "note": "None", + } def test_empty_returns_default(self): - assert _decode_celery_field(None, {}) == {} - assert _decode_celery_field("", []) == [] + assert decode_celery_field(None, {}) == {} + assert decode_celery_field("", []) == [] + assert decode_celery_field("null", {}) == {} + assert decode_celery_field("None", []) == [] + + def test_empty_validates_default(self): + with pytest.raises(ValueError): + decode_celery_field("", {"value": ...}) def test_unparseable_raises(self): with pytest.raises(ValueError): - _decode_celery_field("<>", {}) + decode_celery_field("<>", {}) + + @pytest.mark.parametrize( + "value", + ( + "{'value': ...}", + "{'value': {1, 2}}", + "{'value': b'bytes'}", + '{"value": NaN}', + ), + ) + def test_non_json_values_raise(self, value): + with pytest.raises(ValueError): + decode_celery_field(value, {}) + + def test_truncated_repr_raises(self): + kwargs_repr = saferepr( + {"finding_ids": [str(uuid4()) for _ in range(30)]}, maxlen=1024 + ) + assert "..." in kwargs_repr + + with pytest.raises(ValueError): + decode_celery_field(kwargs_repr, {}) @pytest.mark.django_db @@ -98,6 +152,58 @@ class TestReconcileTaskResults: assert call["kwargs"] == {"tenant_id": str(tenant.id)} assert call["task_id"] != tr.task_id # fresh task id + def test_truncated_kwargs_are_not_reenqueued(self, tenants_fixture): + tenant = tenants_fixture[0] + tr = _orphan_result( + name="tenant-deletion", + kwargs={"tenant_id": str(tenant.id)}, + worker="dead@gone", + created_minutes_ago=60, + ) + tr.task_kwargs = saferepr( + {"finding_ids": [str(uuid4()) for _ in range(30)]}, maxlen=1024 + ) + assert "..." in tr.task_kwargs + tr.save(update_fields=["task_kwargs"]) + p_alive, p_revoke, p_app, mock_task = self._patches(alive=False) + + with ( + p_alive, + p_revoke, + p_app, + patch("tasks.jobs.orphan_recovery._recovery_attempt_count", return_value=1), + ): + result = _reconcile_task_results( + grace_minutes=2, max_attempts=3, window_hours=6, dry_run=False + ) + + assert tr.task_id in result["failed"] + mock_task.apply_async.assert_not_called() + + def test_wrong_kwargs_shape_is_not_reenqueued(self, tenants_fixture): + tr = _orphan_result( + name="tenant-deletion", + kwargs={"tenant_id": str(tenants_fixture[0].id)}, + worker="dead@gone", + created_minutes_ago=60, + ) + tr.task_kwargs = "[]" + tr.save(update_fields=["task_kwargs"]) + p_alive, p_revoke, p_app, mock_task = self._patches(alive=False) + + with ( + p_alive, + p_revoke, + p_app, + patch("tasks.jobs.orphan_recovery._recovery_attempt_count", return_value=1), + ): + result = _reconcile_task_results( + grace_minutes=2, max_attempts=3, window_hours=6, dry_run=False + ) + + assert tr.task_id in result["failed"] + mock_task.apply_async.assert_not_called() + def test_external_integration_task_is_not_reenqueued_by_default( self, tenants_fixture ):