mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-04 02:04:06 +00:00
fix(api): safely decode stored Celery task arguments (#12165)
This commit is contained in:
@@ -0,0 +1 @@
|
||||
`task_args` serialization no longer returns HTTP 500 errors when Celery truncates stored task keyword arguments
|
||||
@@ -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
|
||||
@@ -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"})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user