Arthur 20 часов назад
Родитель
Сommit
5fd7bbe0ae
2 измененных файлов с 61 добавлено и 118 удалено
  1. 43 78
      netbox/circuits/models/circuits.py
  2. 18 40
      netbox/circuits/tests/test_models.py

+ 43 - 78
netbox/circuits/models/circuits.py

@@ -1,7 +1,7 @@
 from django.apps import apps
 from django.apps import apps
 from django.contrib.contenttypes.fields import GenericForeignKey, GenericRelation
 from django.contrib.contenttypes.fields import GenericForeignKey, GenericRelation
 from django.core.exceptions import ValidationError
 from django.core.exceptions import ValidationError
-from django.db import models, router, transaction
+from django.db import models
 from django.urls import reverse
 from django.urls import reverse
 from django.utils.translation import gettext_lazy as _
 from django.utils.translation import gettext_lazy as _
 
 
@@ -394,107 +394,72 @@ class CircuitTermination(
 
 
         circuit_changed = tracking_relevant and self._orig_circuit_id and self._orig_circuit_id != self.circuit_id
         circuit_changed = tracking_relevant and self._orig_circuit_id and self._orig_circuit_id != self.circuit_id
         term_side_changed = tracking_relevant and self._orig_term_side and self._orig_term_side != self.term_side
         term_side_changed = tracking_relevant and self._orig_term_side and self._orig_term_side != self.term_side
-        pointer_moved = is_new or circuit_changed or term_side_changed
 
 
         # Cache objects associated with the terminating object (for filtering)
         # Cache objects associated with the terminating object (for filtering)
         self.cache_related_objects()
         self.cache_related_objects()
 
 
-        if not pointer_moved:
-            super().save(*args, **kwargs)
-            return
-
-        # Collect the pointer writes per circuit, so that a term_side change within one
-        # circuit clears the old side and sets the new one in a single write
-        updates = {}
+        super().save(*args, **kwargs)
 
 
-        # Clear the old termination reference if circuit or term_side changed. Never on insert:
-        # __init__ captured the constructor's values, which may name a live sibling's pointer.
-        if not is_new and (circuit_changed or term_side_changed):
+        # Clear the old termination reference if circuit or term_side changed
+        if circuit_changed or term_side_changed:
             old_termination_name = f'termination_{self._orig_term_side.lower()}'
             old_termination_name = f'termination_{self._orig_term_side.lower()}'
-            updates.setdefault(self._orig_circuit_id, {})[old_termination_name] = None
-
-        # Write the termination row and the pointers which reference it together
-        using = kwargs.get('using') or router.db_for_write(type(self))
-        with transaction.atomic(using=using):
-            super().save(*args, **kwargs)
+            self._set_circuit_terminations(
+                self._orig_circuit_id, {old_termination_name: None}, only_if_references=self.pk
+            )
 
 
+        # Update the cache if this is a new termination or circuit/term_side changed
+        if is_new or circuit_changed or term_side_changed:
             # Update the new circuit's termination reference
             # Update the new circuit's termination reference
             termination_name = f'termination_{self.term_side.lower()}'
             termination_name = f'termination_{self.term_side.lower()}'
-            updates.setdefault(self.circuit_id, {})[termination_name] = self.pk
+            self._set_circuit_terminations(self.circuit_id, {termination_name: self.pk})
 
 
-            # Ordered by PK so concurrent saves take the circuit locks in the same order
-            for circuit_id in sorted(updates):
-                self._set_circuit_terminations(circuit_id, updates[circuit_id], using=using)
+            # Update cached values for subsequent saves
+            self._orig_circuit_id = self.circuit_id
+            self._orig_term_side = self.term_side
 
 
-        # Advanced only once the writes have left the block, so that a rolled-back save is still
-        # pending on retry
-        self._orig_circuit_id = self.circuit_id
-        self._orig_term_side = self.term_side
+    def delete(self, *args, **kwargs):
+        # Clear the circuit's reference before the row goes away, so that the change is recorded
+        # and precedes the DELETE. on_delete=SET_NULL would clear it with an unlogged bulk update.
+        self._set_circuit_terminations(
+            self.circuit_id, {'termination_a': None, 'termination_z': None}, only_if_references=self.pk
+        )
+
+        return super().delete(*args, **kwargs)
+
+    delete.alters_data = True
 
 
     @staticmethod
     @staticmethod
-    def _set_circuit_terminations(circuit_id, fields, using=None, only_if_references=None):
+    def _set_circuit_terminations(circuit_id, fields, only_if_references=None):
         """
         """
-        Set or clear a Circuit's cached `termination_a`/`termination_z` fields. `fields` maps
-        field name to CircuitTermination PK (or None). `only_if_references` restricts the write
-        to fields which currently hold that PK.
-
-        Written via snapshot() + save() rather than a queryset update(), which emits no post_save
-        and so records nothing in the changelog. The Circuit is re-fetched under a lock so that
-        the snapshot reflects a sibling pointer written concurrently; without it, a second writer
-        could record the first writer's pointer as null. no_key avoids blocking the foreign key
-        inserts which reference this circuit.
-
-        Does nothing if the Circuit no longer exists, or if every field already holds its
-        intended value.
+        Set or clear a Circuit's cached `termination_a`/`termination_z` fields, recording the
+        change. `fields` maps field name to CircuitTermination PK (or None). A queryset update()
+        emits no post_save, and so records nothing in the changelog (#23134).
+
+        Args:
+            circuit_id: PK of the Circuit to update
+            fields: Mapping of field name to the value to assign
+            only_if_references: If set, restricts the write to fields which currently hold this PK
         """
         """
-        using = using or router.db_for_write(Circuit)
-
-        # order_by() clears the default ordering, whose JOIN would leave the row unlockable
-        circuit = Circuit.objects.using(using).filter(
-            pk=circuit_id
-        ).order_by().select_for_update(no_key=True).first()
+        # Re-fetched rather than reusing a cached circuit, whose pointers may predate a sibling write
+        circuit = Circuit.objects.filter(pk=circuit_id).first()
         if circuit is None:
         if circuit is None:
             return
             return
 
 
-        def needs_write(field_name, value):
+        updates = {}
+        for field_name, value in fields.items():
             current = getattr(circuit, f'{field_name}_id')
             current = getattr(circuit, f'{field_name}_id')
             if current == value:
             if current == value:
-                return False
-            # Match what on_delete=SET_NULL would have cleared
-            return only_if_references is None or current == only_if_references
-
-        fields = {name: value for name, value in fields.items() if needs_write(name, value)}
-        if not fields:
+                continue
+            if only_if_references is not None and current != only_if_references:
+                continue
+            updates[field_name] = value
+        if not updates:
             return
             return
 
 
         circuit.snapshot()
         circuit.snapshot()
-        for field_name, value in fields.items():
+        for field_name, value in updates.items():
             setattr(circuit, f'{field_name}_id', value)
             setattr(circuit, f'{field_name}_id', value)
-
-        # Saved in full, not with update_fields: the mixin chain also mutates custom_field_data
-        # and the distance fields, which would reach postchange_data but not the database.
-        circuit.save(using=using)
-
-    def delete(self, *args, **kwargs):
-        # Clear the pointer first, so its record precedes this DELETE. Not a pre_delete receiver:
-        # handle_deleted_object connects earlier and would record the DELETE first. (#23134)
-        using = kwargs.get('using') or (args[0] if args else None) or router.db_for_write(type(self))
-        with transaction.atomic(using=using):
-            # Locked before the circuit, matching the order super().save() takes them in
-            CircuitTermination.objects.using(using).filter(
-                pk=self.pk
-            ).order_by().select_for_update().first()
-
-            if self.term_side:
-                self._set_circuit_terminations(
-                    self.circuit_id,
-                    {f'termination_{self.term_side.lower()}': None},
-                    using=using,
-                    only_if_references=self.pk,
-                )
-            return super().delete(*args, **kwargs)
-
-    delete.alters_data = True
+        circuit.save(update_fields=[*updates, 'last_updated'])
 
 
     def cache_related_objects(self):
     def cache_related_objects(self):
         self._provider_network = self._region = self._site_group = self._site = self._location = None
         self._provider_network = self._region = self._site_group = self._site = self._location = None

+ 18 - 40
netbox/circuits/tests/test_models.py

@@ -1,5 +1,4 @@
 import uuid
 import uuid
-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, ValidationError
 from django.core.exceptions import NON_FIELD_ERRORS, ValidationError
@@ -383,7 +382,7 @@ class CircuitTerminationChangeLoggingTestCase(TestCase):
         self.assertEqual(new_changes[0].postchange_data['termination_a'], termination.pk)
         self.assertEqual(new_changes[0].postchange_data['termination_a'], termination.pk)
 
 
     @tag('regression')  # Ref: #23134
     @tag('regression')  # Ref: #23134
-    def test_term_side_change_records_single_circuit_update(self):
+    def test_term_side_change_records_circuit_updates(self):
         termination = self._tracked(lambda: CircuitTermination.objects.create(
         termination = self._tracked(lambda: CircuitTermination.objects.create(
             circuit=self.circuits[0], term_side='A', termination=self.sites[0],
             circuit=self.circuits[0], term_side='A', termination=self.sites[0],
         ))
         ))
@@ -395,13 +394,13 @@ class CircuitTerminationChangeLoggingTestCase(TestCase):
 
 
         self._tracked(_flip)
         self._tracked(_flip)
 
 
-        # Both pointers move within one circuit, so the clear and the set are coalesced
+        # The old pointer is cleared, then the new one is set
         changes = self._circuit_changes(self.circuits[0])
         changes = self._circuit_changes(self.circuits[0])
-        self.assertEqual(changes.count(), 1)
+        self.assertEqual(changes.count(), 2)
         self.assertEqual(changes[0].prechange_data['termination_a'], termination.pk)
         self.assertEqual(changes[0].prechange_data['termination_a'], termination.pk)
-        self.assertIsNone(changes[0].prechange_data['termination_z'])
         self.assertIsNone(changes[0].postchange_data['termination_a'])
         self.assertIsNone(changes[0].postchange_data['termination_a'])
-        self.assertEqual(changes[0].postchange_data['termination_z'], termination.pk)
+        self.assertIsNone(changes[1].prechange_data['termination_z'])
+        self.assertEqual(changes[1].postchange_data['termination_z'], termination.pk)
 
 
     @tag('regression')  # Ref: #23134
     @tag('regression')  # Ref: #23134
     def test_redundant_pointer_write_is_skipped(self):
     def test_redundant_pointer_write_is_skipped(self):
@@ -550,9 +549,10 @@ class CircuitTerminationChangeLoggingTestCase(TestCase):
         self.assertIsNone(self.circuits[0].termination_a_id)
         self.assertIsNone(self.circuits[0].termination_a_id)
         self.assertFalse(self._circuit_changes(self.circuits[0]).exists())
         self.assertFalse(self._circuit_changes(self.circuits[0]).exists())
 
 
-    def test_deletion_leaves_pointer_for_another_termination(self):
-        # An in-memory term_side which diverges from the persisted one must not clear a pointer
-        # belonging to a different termination
+    @tag('regression')  # Ref: #23134
+    def test_deletion_clears_the_pointer_which_references_it(self):
+        # An in-memory term_side which diverges from the persisted one must clear this
+        # termination's own pointer, and leave the one belonging to its sibling alone
         termination_a = self._tracked(lambda: CircuitTermination.objects.create(
         termination_a = self._tracked(lambda: CircuitTermination.objects.create(
             circuit=self.circuits[0], term_side='A', termination=self.sites[0],
             circuit=self.circuits[0], term_side='A', termination=self.sites[0],
         ))
         ))
@@ -560,13 +560,21 @@ class CircuitTerminationChangeLoggingTestCase(TestCase):
             circuit=self.circuits[0], term_side='Z', termination=self.sites[1],
             circuit=self.circuits[0], term_side='Z', termination=self.sites[1],
         ))
         ))
         ObjectChange.objects.all().delete()
         ObjectChange.objects.all().delete()
+        termination_z_pk = termination_z.pk
 
 
         termination_z.term_side = 'A'
         termination_z.term_side = 'A'
         self._tracked(termination_z.delete)
         self._tracked(termination_z.delete)
 
 
         self.circuits[0].refresh_from_db()
         self.circuits[0].refresh_from_db()
         self.assertEqual(self.circuits[0].termination_a_id, termination_a.pk)
         self.assertEqual(self.circuits[0].termination_a_id, termination_a.pk)
-        self.assertFalse(self._circuit_changes(self.circuits[0]).exists())
+        self.assertIsNone(self.circuits[0].termination_z_id)
+
+        changes = self._circuit_changes(self.circuits[0])
+        self.assertEqual(changes.count(), 1)
+        self.assertEqual(changes[0].prechange_data['termination_a'], termination_a.pk)
+        self.assertEqual(changes[0].postchange_data['termination_a'], termination_a.pk)
+        self.assertEqual(changes[0].prechange_data['termination_z'], termination_z_pk)
+        self.assertIsNone(changes[0].postchange_data['termination_z'])
 
 
     def test_circuit_deletion_records_no_pointer_update(self):
     def test_circuit_deletion_records_no_pointer_update(self):
         self._tracked(lambda: CircuitTermination.objects.create(
         self._tracked(lambda: CircuitTermination.objects.create(
@@ -577,33 +585,3 @@ class CircuitTerminationChangeLoggingTestCase(TestCase):
         self._tracked(self.circuits[0].delete)
         self._tracked(self.circuits[0].delete)
 
 
         self.assertFalse(self._circuit_changes(self.circuits[0]).exists())
         self.assertFalse(self._circuit_changes(self.circuits[0]).exists())
-
-    @tag('regression')  # Ref: #23134
-    def test_failed_pointer_write_leaves_the_change_pending(self):
-        # The cached originals must not advance until the pointer writes have succeeded
-        termination = self._tracked(lambda: CircuitTermination.objects.create(
-            circuit=self.circuits[0], term_side='A', termination=self.sites[0],
-        ))
-
-        def _move():
-            termination.circuit = self.circuits[1]
-            termination.save()
-
-        with patch.object(
-            CircuitTermination, '_set_circuit_terminations', side_effect=OSError('boom')
-        ):
-            with self.assertRaises(OSError):
-                self._tracked(_move)
-
-        # The termination row was rolled back along with the pointer writes
-        termination.refresh_from_db()
-        self.assertEqual(termination.circuit, self.circuits[0])
-
-        # A retry still sees the move as pending, so both pointers end up correct
-        termination.circuit = self.circuits[1]
-        self._tracked(termination.save)
-
-        self.circuits[0].refresh_from_db()
-        self.circuits[1].refresh_from_db()
-        self.assertIsNone(self.circuits[0].termination_a_id)
-        self.assertEqual(self.circuits[1].termination_a_id, termination.pk)