Преглед на файлове

fix(api): Pass database alias to transaction rollback in bulk operations

Adds missing `using` parameter to `transaction.set_rollback()` calls in
bulk create, update, and delete operations. Ensures transactions roll
back on the correct database connection when using DATABASE_ROUTERS.

Fixes #23367
Martin Hauser преди 21 часа
родител
ревизия
d9f1b4177a
променени са 2 файла, в които са добавени 107 реда и са изтрити 6 реда
  1. 3 3
      netbox/netbox/api/viewsets/mixins.py
  2. 104 3
      netbox/netbox/tests/test_api.py

+ 3 - 3
netbox/netbox/api/viewsets/mixins.py

@@ -508,7 +508,7 @@ class BulkCreateModelMixin:
                 else:
                 else:
                     created_pks.append(serializer.instance.pk)
                     created_pks.append(serializer.instance.pk)
             if errors:
             if errors:
-                transaction.set_rollback(True)
+                transaction.set_rollback(True, using=using)
         return created_pks, errors, resolve_bulk_error_status(error_statuses)
         return created_pks, errors, resolve_bulk_error_status(error_statuses)
 
 
 
 
@@ -645,7 +645,7 @@ class BulkUpdateModelMixin:
                 else:
                 else:
                     updated_pks.append(obj.pk)
                     updated_pks.append(obj.pk)
             if errors:
             if errors:
-                transaction.set_rollback(True)
+                transaction.set_rollback(True, using=using)
         return updated_pks, errors, resolve_bulk_error_status(error_statuses)
         return updated_pks, errors, resolve_bulk_error_status(error_statuses)
 
 
     def get_bulk_update_serializer_class(self, *, partial=False):
     def get_bulk_update_serializer_class(self, *, partial=False):
@@ -795,7 +795,7 @@ class BulkDestroyModelMixin:
                     errors.append({'id': pk, 'errors': {'__all__': [PERMISSION_DENIED_MESSAGE]}})
                     errors.append({'id': pk, 'errors': {'__all__': [PERMISSION_DENIED_MESSAGE]}})
                     error_statuses.add(status.HTTP_403_FORBIDDEN)
                     error_statuses.add(status.HTTP_403_FORBIDDEN)
             if errors:
             if errors:
-                transaction.set_rollback(True)
+                transaction.set_rollback(True, using=using)
         return errors, total, resolve_bulk_error_status(error_statuses)
         return errors, total, resolve_bulk_error_status(error_statuses)
 
 
 
 

+ 104 - 3
netbox/netbox/tests/test_api.py

@@ -4,13 +4,15 @@ from unittest.mock import patch
 
 
 from django.contrib.contenttypes.models import ContentType
 from django.contrib.contenttypes.models import ContentType
 from django.core.exceptions import NON_FIELD_ERRORS
 from django.core.exceptions import NON_FIELD_ERRORS
+from django.db import DEFAULT_DB_ALIAS, connections
 from django.db.backends.postgresql.psycopg_any import NumericRange
 from django.db.backends.postgresql.psycopg_any import NumericRange
-from django.test import RequestFactory, TestCase
+from django.test import RequestFactory, TestCase, TransactionTestCase, override_settings, tag
 from django.urls import reverse
 from django.urls import reverse
 from django.utils.timezone import now
 from django.utils.timezone import now
 from rest_framework.exceptions import ValidationError
 from rest_framework.exceptions import ValidationError
 from rest_framework.request import Request
 from rest_framework.request import Request
 from rest_framework.settings import api_settings
 from rest_framework.settings import api_settings
+from rest_framework.test import APIClient
 
 
 from core.models import DataFile, DataSource, ObjectType
 from core.models import DataFile, DataSource, ObjectType
 from dcim.api.serializers import RackSerializer
 from dcim.api.serializers import RackSerializer
@@ -22,9 +24,10 @@ from netbox.api.exceptions import QuerySetNotOrdered, SerializerNotFound
 from netbox.api.fields import ContentTypeField, IntegerRangeSerializer, RelatedObjectCountField
 from netbox.api.fields import ContentTypeField, IntegerRangeSerializer, RelatedObjectCountField
 from netbox.api.pagination import NetBoxPagination
 from netbox.api.pagination import NetBoxPagination
 from netbox.api.serializers import ValidatedModelSerializer
 from netbox.api.serializers import ValidatedModelSerializer
-from users.models import Token
+from users.constants import TOKEN_PREFIX
+from users.models import Token, User
 from utilities.api import get_serializer_for_model
 from utilities.api import get_serializer_for_model
-from utilities.testing import APITestCase
+from utilities.testing import APITestCase, create_test_device
 from vpn.api.serializers import L2VPNSerializer
 from vpn.api.serializers import L2VPNSerializer
 
 
 
 
@@ -437,3 +440,101 @@ class ValidatedModelSerializerTestCase(TestCase):
         data = {'template_code': '{# untouched #}'}
         data = {'template_code': '{# untouched #}'}
 
 
         self.assertEqual(serializer.validate(data), data)
         self.assertEqual(serializer.validate(data), data)
+
+
+ROUTED_ALIAS = 'routed'
+
+
+class SiteRouter:
+    """
+    Route Site queries to a second connection, as a plugin's router may.
+    """
+    def db_for_read(self, model, **hints):
+        return ROUTED_ALIAS if model is Site else None
+
+    db_for_write = db_for_read
+
+    def allow_relation(self, obj1, obj2, **hints):
+        return True
+
+
+@override_settings(DATABASE_ROUTERS=[SiteRouter()])
+class BulkOperationRoutingTestCase(TransactionTestCase):
+    """
+    Exercise the bulk operations on a connection which DATABASE_ROUTERS selects in place of the default one, as
+    netbox-branching does for an active branch. Uses TransactionTestCase so that the default connection is not in
+    a transaction while the request is served, as is the case outside of tests.
+
+    Note: TransactionTestCase teardown flushes all tables, which removes rows seeded by data migrations from a
+    --keepdb database (e.g. the dcim.0206 ModuleTypeProfiles). A fresh test database restores them.
+    """
+    client_class = APIClient
+
+    def setUp(self):
+        # A second connection to the test database, standing in for e.g. a branch schema
+        routed = connections[DEFAULT_DB_ALIAS].copy(ROUTED_ALIAS)
+        connections[ROUTED_ALIAS] = routed
+        self.addCleanup(connections.__delitem__, ROUTED_ALIAS)
+        self.addCleanup(routed.close)
+
+        # A superuser, as this case covers transaction handling rather than permission enforcement
+        user = User.objects.create_user(username='testuser', is_superuser=True)
+        token = Token.objects.create(user=user)
+        self.header = {'HTTP_AUTHORIZATION': f'Bearer {TOKEN_PREFIX}{token.key}.{token.token}'}
+        self.url = reverse('dcim-api:site-list')
+
+    @tag('regression')  # Ref: #23367
+    def test_bulk_create_rollback(self):
+        """Roll back a bulk create on the routed connection when one object is invalid."""
+        data = [
+            {'name': 'Site 1', 'slug': 'site-1'},
+            {'name': 'Site 2'},
+        ]
+        with patch('netbox.context_managers.flush_events') as flush_events:
+            response = self.client.post(self.url, data, format='json', **self.header)
+
+        self.assertEqual(response.status_code, 400)
+        self.assertEqual([e['index'] for e in response.data['errors']], [1])
+        self.assertIn('slug', response.data['errors'][0]['errors'])
+        self.assertFalse(Site.objects.exists())
+        flush_events.assert_not_called()
+
+    @tag('regression')  # Ref: #23367
+    def test_bulk_update_rollback(self):
+        """Roll back a bulk update on the routed connection when one object is invalid."""
+        sites = (
+            Site(name='Site 1', slug='site-1'),
+            Site(name='Site 2', slug='site-2'),
+        )
+        Site.objects.bulk_create(sites)
+        data = [
+            {'id': sites[0].pk, 'description': 'Updated'},
+            {'id': sites[1].pk, 'status': 'invalid'},
+        ]
+        with patch('netbox.context_managers.flush_events') as flush_events:
+            response = self.client.patch(self.url, data, format='json', **self.header)
+
+        self.assertEqual(response.status_code, 400)
+        self.assertEqual([e['id'] for e in response.data['errors']], [sites[1].pk])
+        self.assertIn('status', response.data['errors'][0]['errors'])
+        self.assertEqual(Site.objects.get(pk=sites[0].pk).description, '')
+        flush_events.assert_not_called()
+
+    @tag('regression')  # Ref: #23367
+    def test_bulk_delete_rollback(self):
+        """Roll back a bulk delete on the routed connection when one object is protected."""
+        sites = (
+            Site(name='Site 1', slug='site-1'),
+            Site(name='Site 2', slug='site-2'),
+        )
+        Site.objects.bulk_create(sites)
+        # Protect the first Site, as its failed delete clears the event queue before the second is deleted
+        create_test_device('Device 1', site=sites[0])
+        data = [{'id': site.pk} for site in sites]
+        with patch('netbox.context_managers.flush_events') as flush_events:
+            response = self.client.delete(self.url, data, format='json', **self.header)
+
+        self.assertEqual(response.status_code, 409)
+        self.assertEqual([e['id'] for e in response.data['errors']], [sites[0].pk])
+        self.assertEqual(Site.objects.count(), 2)
+        flush_events.assert_not_called()