瀏覽代碼

Applied suggestions from Claude review

Added new method which is called by extras/views.py and extras/api/views.py so
these won't differ.

Added correct tests.
Martin Burggraf 3 周之前
父節點
當前提交
8774410243
共有 5 個文件被更改,包括 180 次插入156 次删除
  1. 38 54
      netbox/extras/api/views.py
  2. 16 0
      netbox/extras/scripts.py
  3. 124 2
      netbox/extras/tests/test_api.py
  4. 0 93
      netbox/extras/tests/test_jobs.py
  5. 2 7
      netbox/extras/views.py

+ 38 - 54
netbox/extras/api/views.py

@@ -4,7 +4,7 @@ from django.utils.translation import gettext_lazy as _
 from drf_spectacular.utils import OpenApiResponse, OpenApiTypes, extend_schema, extend_schema_view
 from drf_spectacular.utils import OpenApiResponse, OpenApiTypes, extend_schema, extend_schema_view
 from rest_framework import status
 from rest_framework import status
 from rest_framework.decorators import action
 from rest_framework.decorators import action
-from rest_framework.exceptions import PermissionDenied
+from rest_framework.exceptions import PermissionDenied, ValidationError
 from rest_framework.generics import RetrieveUpdateDestroyAPIView
 from rest_framework.generics import RetrieveUpdateDestroyAPIView
 from rest_framework.mixins import CreateModelMixin, ListModelMixin, RetrieveModelMixin, UpdateModelMixin
 from rest_framework.mixins import CreateModelMixin, ListModelMixin, RetrieveModelMixin, UpdateModelMixin
 from rest_framework.renderers import JSONRenderer
 from rest_framework.renderers import JSONRenderer
@@ -16,6 +16,7 @@ from core.choices import ManagedFileRootPathChoices
 from extras import filtersets
 from extras import filtersets
 from extras.jobs import ScriptJob
 from extras.jobs import ScriptJob
 from extras.models import *
 from extras.models import *
+from extras.scripts import prepare_script_form
 from netbox.api.authentication import IsAuthenticatedOrLoginNotRequired, TokenWritePermission
 from netbox.api.authentication import IsAuthenticatedOrLoginNotRequired, TokenWritePermission
 from netbox.api.features import SyncedDataMixin
 from netbox.api.features import SyncedDataMixin
 from netbox.api.metadata import ContentTypeMetadata
 from netbox.api.metadata import ContentTypeMetadata
@@ -366,62 +367,45 @@ class ScriptViewSet(ModelViewSet):
         if not any_workers_for_queue('default'):
         if not any_workers_for_queue('default'):
             raise RQWorkerNotRunningException()
             raise RQWorkerNotRunningException()
 
 
-        if input_serializer.is_valid():
-            # Instantiate the script class so we can validate/clean the input via its form.
-            script_class = script.python_class
-            script_instance = script_class()
-
-            # Prepare payload and files
-            payload = input_serializer.validated_data.get('data', {}) or {}
-            files = request.FILES if request else None
-
-            # Validate via the script's form so ObjectVar/MultiObjectVar IDs get converted
-            try:
-                form = script_instance.as_form(data=payload, files=files)
-            except Exception as e:
-                # Defensive: if form construction raises, respond 400 with a helpful message.
-                return Response({'detail': f"Error preparing script form: {e}"}, status=status.HTTP_400_BAD_REQUEST)
-
-            if not form.is_valid():
-                # Return form errors as a 400 so clients get immediate feedback (instead of a failed background job)
-                return Response(form.errors, status=status.HTTP_400_BAD_REQUEST)
-
-            # Use cleaned_data for execution parameters and script variables
-            cleaned = dict(form.cleaned_data)
-
-            # Pop known execution parameters explicitly (do not generically strip _-prefixed names)
-            schedule_at = cleaned.pop('_schedule_at', input_serializer.validated_data.get('schedule_at'))
-            interval = cleaned.pop('_interval', input_serializer.validated_data.get('interval'))
-            notifications = cleaned.pop('_notifications', input_serializer.validated_data.get('notifications'))
-            commit = cleaned.pop(
-                '_commit',
-                input_serializer.validated_data.get('commit', script_instance.commit_default)
-            )
-
-            # Ensure any uploaded files are preserved if not claimed by the form
-            if files:
-                for fname, fobj in files.items():
-                    if fname not in cleaned:
-                        cleaned[fname] = fobj
-
-            # Enqueue the job with cleaned data (model instances/QuerySets where appropriate)
-            ScriptJob.enqueue(
-                instance=script,
-                user=request.user,
-                data=cleaned,
-                request=copy_safe_request(request),
-                commit=commit,
-                job_timeout=script_class.job_timeout,
-                schedule_at=schedule_at,
-                interval=interval,
-                notifications=notifications,
-            )
+        if not input_serializer.is_valid():
+            return Response(input_serializer.errors, status=status.HTTP_400_BAD_REQUEST)
 
 
-            serializer = serializers.ScriptDetailSerializer(script, context={'request': request})
+        validated = input_serializer.validated_data
 
 
-            return Response(serializer.data)
+        payload = validated.get('data')
+        if not isinstance(payload, dict):
+            raise ValidationError(
+                {'data': _('Invalid data payload; expected an object mapping variable names to values.')}
+            )
 
 
-        return Response(input_serializer.errors, status=status.HTTP_400_BAD_REQUEST)
+        script_class = script.python_class
+        if not script_class:
+            raise ValidationError({'script': _('Script class could not be loaded; cannot determine job timeout.')})
+        script_instance = script_class()
+
+        form = prepare_script_form(script_instance, payload, files=request.FILES)
+        if not form.is_valid():
+            # remove internal fields (_commit etc.) from API error message
+            errors = {k: v for k, v in form.errors.items() if not k.startswith('_')}
+            raise ValidationError(errors)
+
+        data = form.cleaned_data.copy()
+        for k in ('_commit', '_schedule_at', '_interval', '_notifications'):
+            data.pop(k, None)
+
+        ScriptJob.enqueue(
+            instance=script,
+            user=request.user,
+            data=data,
+            request=copy_safe_request(request),
+            commit=validated.get('commit'),
+            job_timeout=script_class.job_timeout,
+            schedule_at=validated.get('schedule_at'),
+            interval=validated.get('interval'),
+            notifications=validated.get('notifications'),
+        )
+        serializer = serializers.ScriptDetailSerializer(script, context={'request': request})
+        return Response(serializer.data)
 
 
 
 
 #
 #

+ 16 - 0
netbox/extras/scripts.py

@@ -41,6 +41,7 @@ __all__ = (
     'StringVar',
     'StringVar',
     'TextVar',
     'TextVar',
     'get_module_and_script',
     'get_module_and_script',
+    'prepare_script_form',
 )
 )
 
 
 
 
@@ -655,3 +656,18 @@ def get_module_and_script(module_name, script_name):
     module = ScriptModule.objects.get(file_path=f'{module_name}.py')
     module = ScriptModule.objects.get(file_path=f'{module_name}.py')
     script = module.scripts.get(name=script_name)
     script = module.scripts.get(name=script_name)
     return module, script
     return module, script
+
+
+def prepare_script_form(script_instance, data, files=None):
+    """
+    Build a bound ScriptForm for an already-instantiated Script object, back-filling any
+    declared variable's `default` value into `data` when the caller omitted it.
+
+    Used by both the UI (extras/views.py) and the REST API (extras/api/views.py) so the
+    two entry points share one contract and can't drift apart again.
+    """
+    data = data.copy() if data is not None else {}
+    for name, var in script_instance._get_vars().items():
+        if name not in data and (initial := var.field_attrs.get('initial')) is not None:
+            data[name] = initial
+    return script_instance.as_form(data=data, files=files)

+ 124 - 2
netbox/extras/tests/test_api.py

@@ -12,13 +12,13 @@ from django.urls import reverse
 from django.utils.timezone import make_aware, now
 from django.utils.timezone import make_aware, now
 from rest_framework import status
 from rest_framework import status
 
 
-from core.choices import ManagedFileRootPathChoices
+from core.choices import JobNotificationChoices, ManagedFileRootPathChoices
 from core.events import *
 from core.events import *
 from core.models import DataFile, DataSource, ObjectType
 from core.models import DataFile, DataSource, ObjectType
 from dcim.models import Device, DeviceRole, DeviceType, Location, Manufacturer, Rack, RackRole, Site
 from dcim.models import Device, DeviceRole, DeviceType, Location, Manufacturer, Rack, RackRole, Site
 from extras.choices import *
 from extras.choices import *
 from extras.models import *
 from extras.models import *
-from extras.scripts import BooleanVar, IntegerVar, StringVar
+from extras.scripts import BooleanVar, IntegerVar, MultiObjectVar, ObjectVar, StringVar
 from extras.scripts import Script as PythonClass
 from extras.scripts import Script as PythonClass
 from users.constants import TOKEN_PREFIX
 from users.constants import TOKEN_PREFIX
 from users.models import Group, ObjectPermission, Token, User
 from users.models import Group, ObjectPermission, Token, User
@@ -1444,6 +1444,128 @@ class ScriptTestCase(APITestCase):
             self.TestScriptClass.Meta.scheduling_enabled = original
             self.TestScriptClass.Meta.scheduling_enabled = original
 
 
 
 
+class ScriptRunExecutionTestCase(APITestCase):
+    """
+    Exercises ScriptViewSet.post() end-to-end (real request -> real serializer -> real
+    form -> real ScriptJob.enqueue() call), covering the regressions raised in review of
+    PR #22861: execution parameters must be taken from the validated request rather than
+    the form's own defaults, ObjectVar/MultiObjectVar values must be converted from raw
+    IDs to model instances/querysets, and declared defaults must be back-filled for
+    variables the client omits.
+    """
+
+    class TestScriptClass(PythonClass):
+        class Meta:
+            name = 'Test run script'
+
+        site = ObjectVar(model=Site)
+        sites = MultiObjectVar(model=Site, required=False)
+        label = StringVar(default='hello')
+
+        def run(self, data, commit=True):
+            return 'ok'
+
+    @classmethod
+    def setUpTestData(cls):
+        cls.sites = [
+            Site.objects.create(name=f'Test Site {i}', slug=f'test-site-{i}') for i in range(1, 3)
+        ]
+        with patch.object(ScriptModule, 'sync_classes'):
+            module = ScriptModule.objects.create(
+                file_root=ManagedFileRootPathChoices.SCRIPTS,
+                file_path='run_script.py',
+            )
+        script = Script.objects.create(
+            module=module,
+            name='Test run script',
+            is_executable=True,
+        )
+        cls.url = reverse('extras-api:script-detail', kwargs={'pk': script.pk})
+
+    def setUp(self):
+        super().setUp()
+        self.add_permissions('extras.run_script')
+
+        # Monkey-patch the Script model to return our TestScriptClass above
+        Script.python_class = self.TestScriptClass
+
+        # The script-run endpoint gates on a live RQ worker. Tests run without one, so
+        # bypass the check to exercise validation and the enqueue path.
+        worker_patch = patch('extras.api.views.any_workers_for_queue', return_value=True)
+        worker_patch.start()
+        self.addCleanup(worker_patch.stop)
+
+    @patch('extras.jobs.ScriptJob.enqueue')
+    def test_run_forwards_commit_value(self, mock_enqueue):
+        for commit_value in (True, False):
+            with self.subTest(commit=commit_value):
+                mock_enqueue.reset_mock()
+                payload = {'data': {'site': self.sites[0].pk}, 'commit': commit_value}
+
+                response = self.client.post(self.url, payload, format='json', **self.header)
+
+                self.assertHttpStatus(response, status.HTTP_200_OK)
+                mock_enqueue.assert_called_once()
+                self.assertIs(mock_enqueue.call_args.kwargs['commit'], commit_value)
+
+    @patch('extras.jobs.ScriptJob.enqueue')
+    def test_run_forwards_notifications_value(self, mock_enqueue):
+        # Regression: ScriptForm.clean() overwrites an empty '_notifications' with the
+        # field's own initial, so a client-supplied value never reached ScriptJob.enqueue.
+        payload = {
+            'data': {'site': self.sites[0].pk},
+            'commit': True,
+            'notifications': JobNotificationChoices.NOTIFICATION_NEVER,
+        }
+
+        response = self.client.post(self.url, payload, format='json', **self.header)
+
+        self.assertHttpStatus(response, status.HTTP_200_OK)
+        self.assertEqual(
+            mock_enqueue.call_args.kwargs['notifications'],
+            JobNotificationChoices.NOTIFICATION_NEVER,
+        )
+
+    @patch('extras.jobs.ScriptJob.enqueue')
+    def test_run_converts_objectvar_and_multiobjectvar_ids(self, mock_enqueue):
+        payload = {
+            'data': {
+                'site': self.sites[0].pk,
+                'sites': [site.pk for site in self.sites],
+            },
+            'commit': True,
+        }
+
+        response = self.client.post(self.url, payload, format='json', **self.header)
+
+        self.assertHttpStatus(response, status.HTTP_200_OK)
+        data = mock_enqueue.call_args.kwargs['data']
+        self.assertEqual(data['site'], self.sites[0])
+        self.assertEqual(
+            set(data['sites'].values_list('pk', flat=True)),
+            {site.pk for site in self.sites},
+        )
+
+    @patch('extras.jobs.ScriptJob.enqueue')
+    def test_run_backfills_default_for_omitted_required_var(self, mock_enqueue):
+        # Regression: required vars declaring `default=` were not back-filled before
+        # binding the form on the API path (unlike the UI path), so they 400'd even
+        # though the client legitimately omitted them.
+        payload = {'data': {'site': self.sites[0].pk}, 'commit': True}  # 'label' omitted
+
+        response = self.client.post(self.url, payload, format='json', **self.header)
+
+        self.assertHttpStatus(response, status.HTTP_200_OK)
+        self.assertEqual(mock_enqueue.call_args.kwargs['data']['label'], 'hello')
+
+    def test_run_rejects_non_dict_payload(self):
+        payload = {'data': 'not-a-dict', 'commit': True}
+
+        response = self.client.post(self.url, payload, format='json', **self.header)
+
+        self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
+
+
 class CreatedUpdatedFilterTestCase(APITestCase):
 class CreatedUpdatedFilterTestCase(APITestCase):
 
 
     @classmethod
     @classmethod

+ 0 - 93
netbox/extras/tests/test_jobs.py

@@ -6,10 +6,7 @@ from unittest.mock import MagicMock, patch
 from django.db import DEFAULT_DB_ALIAS
 from django.db import DEFAULT_DB_ALIAS
 from django.test import TestCase
 from django.test import TestCase
 
 
-from dcim.models import Site
 from extras.jobs import ScriptJob
 from extras.jobs import ScriptJob
-from extras.models import Script as ScriptModel
-from extras.scripts import ObjectVar, Script
 from utilities.exceptions import AbortScript
 from utilities.exceptions import AbortScript
 
 
 
 
@@ -341,93 +338,3 @@ class ScriptJobRunTestCase(TestCase):
             runner.run(data={}, commit=False)
             runner.run(data={}, commit=False)
 
 
         self.assertEqual(entered, ['proc_a'])
         self.assertEqual(entered, ['proc_a'])
-
-
-class ScriptJobFormCleaningTestCase(TestCase):
-    def test_run_converts_objectvar_id_to_model_instance(self):
-        # Create a simple target object
-        site = Site.objects.create(name="Test Site", slug="test-site")
-
-        # Real Script subclass that declares an ObjectVar and records received data
-        class TestScript(Script):
-            site = ObjectVar(label='Site', model=Site)
-
-            def __init__(self):
-                super().__init__()
-                self.received = None
-
-            def run(self, data, commit=True):
-                # record what we received and return a sentinel
-                self.received = data
-                return "ok"
-
-        script_instance = TestScript()
-
-        # Make ScriptModel.objects.get() return a stub whose python_class() yields our instance
-        script_model_stub = MagicMock()
-        script_model_stub.python_class.return_value = script_instance
-
-        runner = _make_runner(object_id=1)
-        with patch.object(ScriptModel.objects, 'get', return_value=script_model_stub):
-            # Simulate what the API view does: build the form and use cleaned_data
-            form = script_instance.as_form(data={'site': site.pk}, files=None)
-            assert form.is_valid(), f"test setup: form invalid: {form.errors}"
-            cleaned = dict(form.cleaned_data)
-            # Pop execution params if present (API does this)
-            cleaned.pop('_commit', None)
-            cleaned.pop('_schedule_at', None)
-            cleaned.pop('_interval', None)
-            cleaned.pop('_notifications', None)
-
-            # Pass the cleaned dict to run(), not the nested {"data": {...}} shape
-            runner.run(data=cleaned, request=None, commit=True)
-
-        # Assert the script received a model instance for 'site'
-        self.assertIsNotNone(script_instance.received)
-        self.assertIn('site', script_instance.received)
-        self.assertIsInstance(script_instance.received['site'], Site)
-        self.assertNotIn('_commit', script_instance.received)
-
-    def test_run_merges_request_files_into_data_for_real_script(self):
-        # Create object
-        site = Site.objects.create(name="Test Site 2", slug="test-site-2")
-
-        class TestScript(Script):
-            site = ObjectVar(label='Site', model=Site)
-
-            def __init__(self):
-                super().__init__()
-                self.received = None
-
-            def run(self, data, commit=True):
-                self.received = data
-                return "ok"
-
-        script_instance = TestScript()
-        script_model_stub = MagicMock()
-        script_model_stub.python_class.return_value = script_instance
-
-        # Simulate a request with uploaded files
-        fake_request = MagicMock()
-        fake_request.FILES = {'upload': 'fileobj'}
-        fake_request.id = None
-
-        runner = _make_runner(object_id=1)
-        with patch.object(ScriptModel.objects, 'get', return_value=script_model_stub):
-            # Build the cleaned data as the API would
-            form = script_instance.as_form(data={'site': site.pk}, files=fake_request.FILES)
-            assert form.is_valid(), f"test setup: form invalid: {form.errors}"
-            cleaned = dict(form.cleaned_data)
-
-            # API merges any uploaded files that the form didn't declare
-            for fname, fobj in fake_request.FILES.items():
-                if fname not in cleaned:
-                    cleaned[fname] = fobj
-
-            runner.run(data=cleaned, request=fake_request, commit=True)
-
-        # Assert both file merged and ObjectVar conversion happened
-        self.assertIsNotNone(script_instance.received)
-        self.assertIn('upload', script_instance.received)
-        self.assertEqual(script_instance.received['upload'], 'fileobj')
-        self.assertIsInstance(script_instance.received['site'], Site)

+ 2 - 7
netbox/extras/views.py

@@ -21,6 +21,7 @@ from dcim.models import Device, DeviceRole, Platform
 from extras.choices import LogLevelChoices
 from extras.choices import LogLevelChoices
 from extras.dashboard.forms import DashboardWidgetAddForm, DashboardWidgetForm
 from extras.dashboard.forms import DashboardWidgetAddForm, DashboardWidgetForm
 from extras.dashboard.utils import get_widget_class
 from extras.dashboard.utils import get_widget_class
+from extras.scripts import prepare_script_form
 from extras.utils import SharedObjectViewMixin
 from extras.utils import SharedObjectViewMixin
 from netbox.object_actions import *
 from netbox.object_actions import *
 from netbox.ui import layout
 from netbox.ui import layout
@@ -1730,13 +1731,7 @@ class ScriptView(BaseScriptView):
                 'script': script,
                 'script': script,
             })
             })
 
 
-        # Populate missing variables with their default values, if defined
-        post_data = request.POST.copy()
-        for name, var in script_class._get_vars().items():
-            if name not in post_data and (initial := var.field_attrs.get('initial')) is not None:
-                post_data[name] = initial
-
-        form = script_class.as_form(post_data, request.FILES)
+        form = prepare_script_form(script_class, request.POST, request.FILES)
 
 
         # Allow execution only if RQ worker process is running
         # Allow execution only if RQ worker process is running
         if not any_workers_for_queue('default'):
         if not any_workers_for_queue('default'):