From 5fd7bbe0aed11c43478a2967691d1bf61cdb83b8 Mon Sep 17 00:00:00 2001 From: Arthur Date: Wed, 16 Sep 2026 09:32:50 -0700 Subject: [PATCH] simplify --- netbox/circuits/models/circuits.py | 119 ++++++++++----------------- netbox/circuits/tests/test_models.py | 58 ++++--------- 2 files changed, 60 insertions(+), 117 deletions(-) diff --git a/netbox/circuits/models/circuits.py b/netbox/circuits/models/circuits.py index b0a420571..b8d7dce2b 100644 --- a/netbox/circuits/models/circuits.py +++ b/netbox/circuits/models/circuits.py @@ -1,7 +1,7 @@ from django.apps import apps from django.contrib.contenttypes.fields import GenericForeignKey, GenericRelation 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.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 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) self.cache_related_objects() - if not pointer_moved: - super().save(*args, **kwargs) - return + super().save(*args, **kwargs) - # 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 = {} - - # 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()}' - 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 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 - 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. + 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). - 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. + 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: return - def needs_write(field_name, value): + updates = {} + for field_name, value in fields.items(): current = getattr(circuit, f'{field_name}_id') 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 circuit.snapshot() - for field_name, value in fields.items(): + for field_name, value in updates.items(): 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): self._provider_network = self._region = self._site_group = self._site = self._location = None diff --git a/netbox/circuits/tests/test_models.py b/netbox/circuits/tests/test_models.py index 9a669f342..0c3c27193 100644 --- a/netbox/circuits/tests/test_models.py +++ b/netbox/circuits/tests/test_models.py @@ -1,5 +1,4 @@ import uuid -from unittest.mock import patch from django.contrib.contenttypes.models import ContentType 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) @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( circuit=self.circuits[0], term_side='A', termination=self.sites[0], )) @@ -395,13 +394,13 @@ class CircuitTerminationChangeLoggingTestCase(TestCase): 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]) - self.assertEqual(changes.count(), 1) + self.assertEqual(changes.count(), 2) 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.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 def test_redundant_pointer_write_is_skipped(self): @@ -550,9 +549,10 @@ class CircuitTerminationChangeLoggingTestCase(TestCase): self.assertIsNone(self.circuits[0].termination_a_id) 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( 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], )) ObjectChange.objects.all().delete() + termination_z_pk = termination_z.pk termination_z.term_side = 'A' self._tracked(termination_z.delete) self.circuits[0].refresh_from_db() 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): self._tracked(lambda: CircuitTermination.objects.create( @@ -577,33 +585,3 @@ class CircuitTerminationChangeLoggingTestCase(TestCase): self._tracked(self.circuits[0].delete) 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)