Sfoglia il codice sorgente

Merge pull request #22646 from netbox-community/20054-bulk-error-correlation

Closes #20054: Return per-object error details for failed bulk operations
bctiemann 1 mese fa
parent
commit
6068f41787

+ 69 - 0
netbox/dcim/tests/test_api.py

@@ -151,6 +151,9 @@ class SiteTestCase(APIViewTestCases.APIViewTestCase):
     bulk_update_data = {
     bulk_update_data = {
         'status': 'planned',
         'status': 'planned',
     }
     }
+    bulk_update_invalid_data = {
+        'status': 'not-a-valid-status',
+    }
     graphql_filter_tests = (
     graphql_filter_tests = (
         GraphQLFilterTest(
         GraphQLFilterTest(
             name='tenant__name__exact',
             name='tenant__name__exact',
@@ -467,6 +470,37 @@ class SiteTestCase(APIViewTestCases.APIViewTestCase):
         response = self.client.patch(url, data, format='json', **self.header)
         response = self.client.patch(url, data, format='json', **self.header)
         self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
         self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
 
 
+    def test_bulk_delete_objects_protected(self):
+        """
+        DELETE a set of objects where one has a protected FK dependency. Verify the structured
+        per-object error response and that no objects are deleted (atomic rollback).
+        """
+        obj_perm = ObjectPermission(name='Test permission', actions=['delete'])
+        obj_perm.save()
+        obj_perm.users.add(self.user)
+        obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
+
+        # Site 1 has no dependent Device; Site 2 gets one (Device FK is on_delete=PROTECT)
+        site1 = Site.objects.get(slug='site-1')
+        site2 = Site.objects.get(slug='site-2')
+        create_test_device('Protected Device', site=site2)
+
+        data = [{'id': site1.pk}, {'id': site2.pk}]
+        response = self.client.delete(self._get_list_url(), data, format='json', **self.header)
+
+        self.assertHttpStatus(response, status.HTTP_409_CONFLICT)
+        self.assertIn('detail', response.data)
+        self.assertIn('errors', response.data)
+        self.assertEqual(len(response.data['errors']), 1)
+
+        # Site 2 (has Device) should be the only entry, since Site 1 succeeded
+        self.assertEqual(response.data['errors'][0]['id'], site2.pk)
+        self.assertIn('errors', response.data['errors'][0])
+
+        # Verify that no sites were actually deleted (transaction rolled back)
+        self.assertTrue(Site.objects.filter(pk=site1.pk).exists(), 'Site 1 should not have been deleted')
+        self.assertTrue(Site.objects.filter(pk=site2.pk).exists(), 'Site 2 should not have been deleted')
+
 
 
 class LocationTestCase(APIViewTestCases.APIViewTestCase):
 class LocationTestCase(APIViewTestCases.APIViewTestCase):
     model = Location
     model = Location
@@ -2207,6 +2241,41 @@ class DeviceTestCase(APIViewTestCases.APIViewTestCase):
         response = self.client.post(url, {'config_template_id': override_template.pk}, format='json', **self.header)
         response = self.client.post(url, {'config_template_id': override_template.pk}, format='json', **self.header)
         self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
         self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
 
 
+    def test_bulk_create_objects_validation_error(self):
+        """
+        POST a set of Device objects where the first passes and the second fails validation.
+        DeviceViewSet uses SequentialBulkCreatesMixin, so the response should report only the
+        failed object, and no objects should be created despite the first item passing
+        (atomic rollback).
+        """
+        obj_perm = ObjectPermission(name='Test permission', actions=['add'])
+        obj_perm.save()
+        obj_perm.users.add(self.user)
+        obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
+
+        self.add_related_view_permissions(self.create_data[0])
+
+        initial_count = self._get_queryset().count()
+        # First item is valid; second is empty (missing required fields) and will fail
+        response = self.client.post(
+            self._get_list_url(),
+            [self.create_data[0], {}],
+            format='json',
+            **self.header,
+        )
+
+        self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
+        self.assertEqual(
+            self._get_queryset().count(), initial_count,
+            'No objects should be created when any sibling fails validation',
+        )
+        self.assertIn('detail', response.data)
+        self.assertIn('errors', response.data)
+        self.assertEqual(len(response.data['errors']), 1)
+        # Second item failed validation — first item succeeded so it's omitted
+        self.assertEqual(response.data['errors'][0]['index'], 1)
+        self.assertIn('errors', response.data['errors'][0])
+
 
 
 class ModuleTestCase(APIViewTestCases.APIViewTestCase):
 class ModuleTestCase(APIViewTestCases.APIViewTestCase):
     model = Module
     model = Module

+ 96 - 16
netbox/netbox/api/viewsets/mixins.py

@@ -1,5 +1,6 @@
 from django.core.exceptions import ObjectDoesNotExist
 from django.core.exceptions import ObjectDoesNotExist
 from django.db import router, transaction
 from django.db import router, transaction
+from django.db.models import ProtectedError, RestrictedError
 from django.http import Http404
 from django.http import Http404
 from django.utils.translation import gettext_lazy as _
 from django.utils.translation import gettext_lazy as _
 from rest_framework import status
 from rest_framework import status
@@ -164,21 +165,44 @@ class SequentialBulkCreatesMixin:
         if (response := handle_background(request, 'create')) is not None:
         if (response := handle_background(request, 'create')) is not None:
             return response
             return response
 
 
+        # Create objects sequentially so each validation sees the state left by prior creates
+        # (e.g. rack space checks). Collect per-object errors instead of failing on the first.
+        errors = []
+        return_data = []
         with transaction.atomic(using=router.db_for_write(self.queryset.model)):
         with transaction.atomic(using=router.db_for_write(self.queryset.model)):
             if not isinstance(request.data, list):
             if not isinstance(request.data, list):
                 # Creating a single object
                 # Creating a single object
                 return super().create(request, *args, **kwargs)
                 return super().create(request, *args, **kwargs)
 
 
-            return_data = []
-            for data in request.data:
+            total = len(request.data)
+            for i, data in enumerate(request.data):
                 serializer = self.get_serializer(data=data)
                 serializer = self.get_serializer(data=data)
-                serializer.is_valid(raise_exception=True)
-                self.perform_create(serializer)
-                return_data.append(serializer.data)
-
-            headers = self.get_success_headers(serializer.data)
+                if serializer.is_valid():
+                    # Provisionally create even when a prior item failed, so subsequent
+                    # cross-object validators (e.g. rack space checks) see a realistic state.
+                    # All creates are rolled back together if any item in the batch fails.
+                    self.perform_create(serializer)
+                    return_data.append(serializer.data)
+                else:
+                    errors.append({'index': i, 'errors': serializer.errors})
+
+            if errors:
+                transaction.set_rollback(True)
+
+        if errors:
+            return Response(
+                {
+                    'detail': _('{failed_count} of {total} objects failed validation.').format(
+                        failed_count=len(errors),
+                        total=total,
+                    ),
+                    'errors': errors,
+                },
+                status=status.HTTP_400_BAD_REQUEST,
+            )
 
 
-            return Response(return_data, status=status.HTTP_201_CREATED, headers=headers)
+        headers = self.get_success_headers(return_data[-1]) if return_data else {}
+        return Response(return_data, status=status.HTTP_201_CREATED, headers=headers)
 
 
 
 
 class BulkUpdateModelMixin:
 class BulkUpdateModelMixin:
@@ -226,7 +250,19 @@ class BulkUpdateModelMixin:
             obj.pop('id'): obj for obj in request.data
             obj.pop('id'): obj for obj in request.data
         }
         }
 
 
-        object_pks = self.perform_bulk_update(qs, update_data, partial=partial)
+        object_pks, errors = self.perform_bulk_update(qs, update_data, partial=partial)
+
+        if errors:
+            return Response(
+                {
+                    'detail': _('{failed_count} of {total} objects failed validation.').format(
+                        failed_count=len(errors),
+                        total=len(object_pks) + len(errors),
+                    ),
+                    'errors': errors,
+                },
+                status=status.HTTP_400_BAD_REQUEST,
+            )
 
 
         # Prefetch related objects for all updated instances
         # Prefetch related objects for all updated instances
         qs = self.get_queryset().filter(pk__in=object_pks)
         qs = self.get_queryset().filter(pk__in=object_pks)
@@ -236,17 +272,24 @@ class BulkUpdateModelMixin:
 
 
     def perform_bulk_update(self, objects, update_data, partial):
     def perform_bulk_update(self, objects, update_data, partial):
         updated_pks = []
         updated_pks = []
+        errors = []
         with transaction.atomic(using=router.db_for_write(self.queryset.model)):
         with transaction.atomic(using=router.db_for_write(self.queryset.model)):
+            # Validate and save each object in turn so subsequent validations see the DB
+            # state left by prior saves (e.g. two items renamed to the same name: the second
+            # will fail validation rather than raising an integrity error on save).
             for obj in objects:
             for obj in objects:
                 data = update_data.get(obj.id)
                 data = update_data.get(obj.id)
                 if hasattr(obj, 'snapshot'):
                 if hasattr(obj, 'snapshot'):
                     obj.snapshot()
                     obj.snapshot()
                 serializer = self.get_serializer(obj, data=data, partial=partial)
                 serializer = self.get_serializer(obj, data=data, partial=partial)
-                serializer.is_valid(raise_exception=True)
-                self.perform_update(serializer)
-                updated_pks.append(obj.pk)
-
-        return updated_pks
+                if serializer.is_valid():
+                    self.perform_update(serializer)
+                    updated_pks.append(obj.pk)
+                else:
+                    errors.append({'id': obj.pk, 'errors': serializer.errors})
+            if errors:
+                transaction.set_rollback(True)
+        return updated_pks, errors
 
 
     def get_bulk_update_serializer_class(self, *, partial=False):
     def get_bulk_update_serializer_class(self, *, partial=False):
         return get_bulk_update_serializer_class(
         return get_bulk_update_serializer_class(
@@ -302,18 +345,55 @@ class BulkDestroyModelMixin:
             o['id']: o.get('changelog_message') for o in serializer.validated_data
             o['id']: o.get('changelog_message') for o in serializer.validated_data
         }
         }
 
 
-        self.perform_bulk_destroy(qs, changelog_messages)
+        errors, total = self.perform_bulk_destroy(qs, changelog_messages)
+
+        if errors:
+            return Response(
+                {
+                    'detail': _('{failed_count} of {total} objects could not be deleted.').format(
+                        failed_count=len(errors),
+                        total=total,
+                    ),
+                    'errors': errors,
+                },
+                status=status.HTTP_409_CONFLICT,
+            )
 
 
         return Response(status=status.HTTP_204_NO_CONTENT)
         return Response(status=status.HTTP_204_NO_CONTENT)
 
 
     def perform_bulk_destroy(self, objects, changelog_messages=None):
     def perform_bulk_destroy(self, objects, changelog_messages=None):
         changelog_messages = changelog_messages or {}
         changelog_messages = changelog_messages or {}
+        errors = []
+        total = 0
         with transaction.atomic(using=router.db_for_write(self.queryset.model)):
         with transaction.atomic(using=router.db_for_write(self.queryset.model)):
             for obj in objects:
             for obj in objects:
+                total += 1
                 if hasattr(obj, 'snapshot'):
                 if hasattr(obj, 'snapshot'):
                     obj.snapshot()
                     obj.snapshot()
                 obj._changelog_message = changelog_messages.get(obj.pk)
                 obj._changelog_message = changelog_messages.get(obj.pk)
-                self.perform_destroy(obj)
+                pk = obj.pk  # Django sets obj.pk = None after deletion; capture it first
+                try:
+                    self.perform_destroy(obj)
+                except (ProtectedError, RestrictedError) as e:
+                    protected = list(
+                        e.protected_objects if isinstance(e, ProtectedError) else e.restricted_objects
+                    )
+                    # Report only the count, not names or PKs, to keep each per-object error
+                    # entry small in a batch response. Note: the single-object delete endpoint
+                    # (NetBoxModelViewSet.dispatch()) does include names and PKs of dependent
+                    # objects, so this is not a hard security boundary — just a narrower
+                    # response shape for the bulk case.
+                    errors.append({
+                        'id': pk,
+                        'errors': {
+                            '__all__': _(
+                                'Unable to delete: {n} dependent object(s) prevent deletion.'
+                            ).format(n=len(protected)),
+                        },
+                    })
+            if errors:
+                transaction.set_rollback(True)
+        return errors, total
 
 
 
 
 class ObjectValidationMixin:
 class ObjectValidationMixin:

+ 46 - 0
netbox/utilities/testing/api.py

@@ -395,6 +395,7 @@ class APIViewTestCases:
     class UpdateObjectViewTestCase(APITestCase):
     class UpdateObjectViewTestCase(APITestCase):
         update_data = {}
         update_data = {}
         bulk_update_data = None
         bulk_update_data = None
+        bulk_update_invalid_data = None
         validation_excluded_fields = []
         validation_excluded_fields = []
 
 
         def test_update_object_without_permission(self):
         def test_update_object_without_permission(self):
@@ -546,6 +547,51 @@ class APIViewTestCases:
                     self.assertObjectChange(oc, action=ObjectChangeActionChoices.ACTION_UPDATE,
                     self.assertObjectChange(oc, action=ObjectChangeActionChoices.ACTION_UPDATE,
                         message=changelog_message)
                         message=changelog_message)
 
 
+        def test_bulk_update_objects_validation_error(self):
+            """
+            PATCH a set of objects where one fails validation. Verify the structured per-object error
+            response and that no objects are modified (atomic rollback).
+            """
+            if self.bulk_update_data is None or self.bulk_update_invalid_data is None:
+                self.skipTest('Bulk update data not set')
+
+            obj_perm = ObjectPermission(name='Test permission', actions=['change'])
+            obj_perm.save()
+            obj_perm.users.add(self.user)
+            obj_perm.object_types.add(ObjectType.objects.get_for_model(self.model))
+
+            id_list = list(self._get_queryset().values_list('id', flat=True)[:2])
+            self.assertEqual(len(id_list), 2, 'Insufficient number of objects to test bulk update validation error')
+
+            # First object: valid data; second: invalid data that must fail validation
+            data = [
+                {'id': id_list[0], **self.bulk_update_data},
+                {'id': id_list[1], **self.bulk_update_invalid_data},
+            ]
+
+            # Snapshot field values before the request so we can verify atomicity afterward
+            instance0_before = self._get_queryset().get(pk=id_list[0])
+
+            response = self.client.patch(self._get_list_url(), data, format='json', **self.header)
+
+            self.assertHttpStatus(response, status.HTTP_400_BAD_REQUEST)
+            self.assertIn('detail', response.data)
+            self.assertIn('errors', response.data)
+            self.assertEqual(len(response.data['errors']), 1)
+            self.assertEqual(response.data['errors'][0]['id'], id_list[1])
+            self.assertIn('errors', response.data['errors'][0])
+
+            # Verify atomicity: object 0 passed validation but must not have been modified
+            instance0_after = self._get_queryset().get(pk=id_list[0])
+            for field in self.bulk_update_data:
+                if field in ('changelog_message', 'add_tags', 'remove_tags'):
+                    continue
+                self.assertEqual(
+                    getattr(instance0_after, field, None),
+                    getattr(instance0_before, field, None),
+                    f'Field {field!r} of object {id_list[0]} was modified — atomic rollback may be broken',
+                )
+
     class DeleteObjectViewTestCase(APITestCase):
     class DeleteObjectViewTestCase(APITestCase):
 
 
         def test_delete_object_without_permission(self):
         def test_delete_object_without_permission(self):