fix(api): safely decode stored Celery task arguments (#12165)

This commit is contained in:
Josema Camacho
2026-07-30 10:30:19 +02:00
committed by GitHub
parent ecf7ec8e85
commit f8be9afa7c
6 changed files with 233 additions and 47 deletions
@@ -0,0 +1 @@
`task_args` serialization no longer returns HTTP 500 errors when Celery truncates stored task keyword arguments
+40
View File
@@ -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
+50
View File
@@ -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"})
+22 -6
View File
@@ -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
@@ -621,13 +626,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
+7 -34
View File
@@ -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(
@@ -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("<<not a literal>>", {})
decode_celery_field("<<not a literal>>", {})
@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
):