Explorar el Código

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 hace 1 día
padre
commit
d9f1b4177a
Se han modificado 2 ficheros con 107 adiciones y 6 borrados
  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:
                     created_pks.append(serializer.instance.pk)
             if errors:
-                transaction.set_rollback(True)
+                transaction.set_rollback(True, using=using)
         return created_pks, errors, resolve_bulk_error_status(error_statuses)
 
 
@@ -645,7 +645,7 @@ class BulkUpdateModelMixin:
                 else:
                     updated_pks.append(obj.pk)
             if errors:
-                transaction.set_rollback(True)
+                transaction.set_rollback(True, using=using)
         return updated_pks, errors, resolve_bulk_error_status(error_statuses)
 
     def get_bulk_update_serializer_class(self, *, partial=False):
@@ -795,7 +795,7 @@ class BulkDestroyModelMixin:
                     errors.append({'id': pk, 'errors': {'__all__': [PERMISSION_DENIED_MESSAGE]}})
                     error_statuses.add(status.HTTP_403_FORBIDDEN)
             if errors:
-                transaction.set_rollback(True)
+                transaction.set_rollback(True, using=using)
         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.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.test import RequestFactory, TestCase
+from django.test import RequestFactory, TestCase, TransactionTestCase, override_settings, tag
 from django.urls import reverse
 from django.utils.timezone import now
 from rest_framework.exceptions import ValidationError
 from rest_framework.request import Request
 from rest_framework.settings import api_settings
+from rest_framework.test import APIClient
 
 from core.models import DataFile, DataSource, ObjectType
 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.pagination import NetBoxPagination
 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.testing import APITestCase
+from utilities.testing import APITestCase, create_test_device
 from vpn.api.serializers import L2VPNSerializer
 
 
@@ -437,3 +440,101 @@ class ValidatedModelSerializerTestCase(TestCase):
         data = {'template_code': '{# untouched #}'}
 
         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()