From f64bf0b217dd5e150ee42eeed2268c342ac4d1ef Mon Sep 17 00:00:00 2001 From: Martin Hauser Date: Mon, 31 Aug 2026 15:51:27 +0200 Subject: [PATCH] fix(models): Normalize update_fields to prevent iterable consumption Introduces normalize_update_fields() utility to materialize one-shot iterables like generators into frozensets, preventing bugs in save() overrides that perform membership tests. Fixes channelization cascades, module moves, and ltree parent tracking when using generator expressions. Fixes #23074 --- netbox/dcim/models/mixins.py | 3 +- netbox/dcim/models/modules.py | 5 ++- netbox/dcim/tests/test_channelization.py | 15 +++++++++ netbox/dcim/tests/test_module_moves.py | 14 ++++++++ netbox/netbox/models/ltree.py | 3 +- netbox/utilities/data.py | 14 ++++++++ netbox/utilities/tests/test_data.py | 31 +++++++++++++++++ netbox/utilities/tests/test_ltree.py | 43 +++++++++++++++++++++--- 8 files changed, 120 insertions(+), 8 deletions(-) diff --git a/netbox/dcim/models/mixins.py b/netbox/dcim/models/mixins.py index adbbdf6e0..dd4697051 100644 --- a/netbox/dcim/models/mixins.py +++ b/netbox/dcim/models/mixins.py @@ -14,6 +14,7 @@ from utilities.conversion import ( to_liters_per_minute, to_millimeters, ) +from utilities.data import normalize_update_fields __all__ = ( 'CachedScopeMixin', @@ -276,7 +277,7 @@ class InterfaceChannelRenameMixin: self._original_channels = self.__dict__.get('channels') def save(self, *args, **kwargs): - update_fields = kwargs.get('update_fields') + update_fields = normalize_update_fields(kwargs) # A save() whose update_fields excludes 'name'/'channels' won't actually persist that attribute, so the # cascade decision below can't treat self.name/self.channels as current in that case -- fall back to the # last known persisted value instead. Without this, e.g. clearing self.channels in memory and saving diff --git a/netbox/dcim/models/modules.py b/netbox/dcim/models/modules.py index 97dc117c3..02b3148c7 100644 --- a/netbox/dcim/models/modules.py +++ b/netbox/dcim/models/modules.py @@ -15,6 +15,7 @@ from extras.models import CustomField from netbox.models import PrimaryModel from netbox.models.features import ImageAttachmentsMixin from netbox.models.mixins import WeightMixin +from utilities.data import normalize_update_fields from utilities.exceptions import AbortRequest from utilities.fields import ColorField, CounterCacheField from utilities.jsonschema import validate_schema @@ -505,11 +506,13 @@ class Module(TrackingModelMixin, PrimaryModel): module = module_module_bay.module if module_module_bay else None def save(self, *args, **kwargs): + # Normalize before the pk branch so _save_new() and _save_existing() forward a re-iterable value. + update_fields = normalize_update_fields(kwargs) + if self.pk is None: self._save_new(*args, **kwargs) return - update_fields = kwargs.get('update_fields') placement_fields = {'device', 'device_id', 'module_bay', 'module_bay_id'} if update_fields is not None and placement_fields.isdisjoint(update_fields): # Placement columns cannot be written by this save, so no move can occur. diff --git a/netbox/dcim/tests/test_channelization.py b/netbox/dcim/tests/test_channelization.py index d2806df33..86a1d26dc 100644 --- a/netbox/dcim/tests/test_channelization.py +++ b/netbox/dcim/tests/test_channelization.py @@ -854,6 +854,21 @@ class ChannelizedInterfaceTestCase(TestCase): self.assertEqual(self.parent.name, 'et0') # Not persisted self.assertEqual(child.name, 'et0:1') # Not cascaded + def test_generator_update_fields_cascades_rename(self): + # A one-shot iterable naming 'name' must still persist the rename and cascade it. + child = Interface.objects.create( + device=self.device, name='et0:1', type=InterfaceTypeChoices.TYPE_CHANNEL, parent=self.parent, channel_id=1 + ) + + self.parent.name = 'et1' + with self.captureOnCommitCallbacks(execute=True): + self.parent.save(update_fields=(field for field in ('name',))) + + self.parent.refresh_from_db() + child.refresh_from_db() + self.assertEqual(self.parent.name, 'et1') + self.assertEqual(child.name, 'et1:1') + def test_update_fields_excluding_name_does_not_desync_later_full_rename(self): # A later full save() must still correctly cascade, proving the earlier partial save didn't refresh # _original_name to its unpersisted in-memory value. diff --git a/netbox/dcim/tests/test_module_moves.py b/netbox/dcim/tests/test_module_moves.py index f8d94b028..75de57b58 100644 --- a/netbox/dcim/tests/test_module_moves.py +++ b/netbox/dcim/tests/test_module_moves.py @@ -529,6 +529,20 @@ class ModuleMoveSaveContractTestCase(TestCase): self.module.refresh_from_db() self.assertEqual(self.module.module_bay, self.bay_b) + def test_generator_update_fields_excluding_module_bay_saves_without_moving(self): + self.module.module_bay = self.bay_b + self.module.serial = 'ABC123' + self.module.save(update_fields=(field for field in ('serial',))) + self.module.refresh_from_db() + self.assertEqual(self.module.serial, 'ABC123') + self.assertEqual(self.module.module_bay, self.bay_a) + + def test_generator_update_fields_including_placement_moves(self): + self.module.module_bay = self.bay_b + self.module.save(update_fields=(field for field in ('device', 'module_bay'))) + self.module.refresh_from_db() + self.assertEqual(self.module.module_bay, self.bay_b) + def test_deadlock_is_translated_to_abort_request(self): deadlock = OperationalError('Simulated deadlock error') deadlock.__cause__ = type('FakeDeadlock', (Exception,), {'sqlstate': '40P01'})() diff --git a/netbox/netbox/models/ltree.py b/netbox/netbox/models/ltree.py index 270fc0ebe..b86b5a986 100644 --- a/netbox/netbox/models/ltree.py +++ b/netbox/netbox/models/ltree.py @@ -23,6 +23,7 @@ from django.utils.translation import gettext_lazy as _ # finishes initializing, so a package-relative import would fail the attribute lookup # on the partially-initialized `netbox.models` package (circular import). from netbox.models.lookups import Ancestor, AncestorOrEqual, Descendant, DescendantOrEqual +from utilities.data import normalize_update_fields from utilities.querysets import RestrictedQuerySet __all__ = ( @@ -387,7 +388,7 @@ class LtreeModel(models.Model, metaclass=LtreeModelBase): # the new parent_id, so the trigger does not fire and _loaded_parent_id # must not advance — otherwise a subsequent full save() would mis-detect # the (real) parent change as already-applied and leave path stale. - update_fields = kwargs.get('update_fields') + update_fields = normalize_update_fields(kwargs) parent_written = update_fields is None or 'parent' in update_fields or 'parent_id' in update_fields parent_changed = (not is_insert) and parent_written and self.parent_id != self._loaded_parent_id diff --git a/netbox/utilities/data.py b/netbox/utilities/data.py index 4cc486be8..fe9cf668c 100644 --- a/netbox/utilities/data.py +++ b/netbox/utilities/data.py @@ -14,6 +14,7 @@ __all__ = ( 'get_config_value_ci', 'get_inclusive_integer_range_bounds', 'normalize_integer_range', + 'normalize_update_fields', 'ranges_to_string', 'ranges_to_string_list', 'resolve_attr_path', @@ -115,6 +116,19 @@ def deep_compare_dict(source_dict, destination_dict, exclude=tuple()): return added, removed +def normalize_update_fields(kwargs): + """ + Replace `kwargs['update_fields']` with a frozenset and return it, so a save() override can + run membership tests without consuming a one-shot iterable. `None` and an absent key are + left alone. + """ + update_fields = kwargs.get('update_fields') + if update_fields is not None: + update_fields = frozenset(update_fields) + kwargs['update_fields'] = update_fields + return update_fields + + # # Array utilities # diff --git a/netbox/utilities/tests/test_data.py b/netbox/utilities/tests/test_data.py index f557bc79e..b969fd717 100644 --- a/netbox/utilities/tests/test_data.py +++ b/netbox/utilities/tests/test_data.py @@ -7,6 +7,7 @@ from utilities.data import ( get_config_value_ci, get_inclusive_integer_range_bounds, normalize_integer_range, + normalize_update_fields, ranges_to_string, ranges_to_string_list, string_to_ranges, @@ -229,3 +230,33 @@ class GetConfigValueCITestCase(TestCase): def test_empty_dict(self): self.assertIsNone(get_config_value_ci({}, 'any.key')) self.assertEqual(get_config_value_ci({}, 'any.key', default=[]), []) + + +class NormalizeUpdateFieldsTestCase(TestCase): + + def test_none_is_passed_through(self): + kwargs = {'update_fields': None} + self.assertIsNone(normalize_update_fields(kwargs)) + self.assertIsNone(kwargs['update_fields']) + + def test_absent_key_is_not_added(self): + kwargs = {} + self.assertIsNone(normalize_update_fields(kwargs)) + self.assertNotIn('update_fields', kwargs) + + def test_generator_is_materialized_in_place(self): + kwargs = {'update_fields': (field for field in ('name', 'description'))} + update_fields = normalize_update_fields(kwargs) + + self.assertEqual(update_fields, frozenset({'name', 'description'})) + self.assertEqual(kwargs['update_fields'], frozenset({'name', 'description'})) + + def test_empty_generator_normalizes_to_empty_frozenset(self): + kwargs = {'update_fields': (field for field in ())} + self.assertEqual(normalize_update_fields(kwargs), frozenset()) + self.assertEqual(kwargs['update_fields'], frozenset()) + + def test_list_is_normalized(self): + kwargs = {'update_fields': ['name']} + self.assertEqual(normalize_update_fields(kwargs), frozenset({'name'})) + self.assertEqual(kwargs['update_fields'], frozenset({'name'})) diff --git a/netbox/utilities/tests/test_ltree.py b/netbox/utilities/tests/test_ltree.py index d61e5b30a..1c4d9a973 100644 --- a/netbox/utilities/tests/test_ltree.py +++ b/netbox/utilities/tests/test_ltree.py @@ -354,13 +354,46 @@ class AddRelatedCountTests(TestCase): class SaveUpdateFieldsTests(TestCase): - """ - Regression: when save(update_fields=...) excludes parent, _loaded_parent_id - must not advance, otherwise a subsequent full save() will mis-detect the - parent change as already-applied and leave path stale in memory. - """ + """Verify partial saves honor Django's update_fields contract.""" + + def test_generator_update_fields_persists_named_field(self): + """A one-shot iterable naming only 'name' still reaches the database.""" + region = Region.objects.create(name='Original', slug='obj-uf-gen') + + region.name = 'Updated' + region.save(update_fields=(field for field in ('name',))) + + db = Region.objects.values('name', 'sort_path').get(pk=region.pk) + self.assertEqual(db['name'], 'Updated') + self.assertEqual(db['sort_path'], 'Updated') + self.assertEqual(region.sort_path, db['sort_path']) + + def test_generator_update_fields_reparent_refreshes_path(self): + """A one-shot iterable naming 'parent' persists the move and refreshes path.""" + parent = Region.objects.create(name='P', slug='p-uf-gen') + child = Region.objects.create(name='C', slug='c-uf-gen') + + child.parent = parent + child.save(update_fields=(field for field in ('parent',))) + + db_path = Region.objects.values_list('path', flat=True).get(pk=child.pk) + self.assertEqual(db_path, _path(parent.pk, child.pk)) + self.assertEqual(child.path, db_path) + + def test_empty_generator_update_fields_skips_save(self): + """An empty iterable is a no-op save, per Django's update_fields contract.""" + region = Region.objects.create(name='Original', slug='empty-uf-gen') + + region.name = 'Updated' + region.save(update_fields=(field for field in ())) + + self.assertEqual(Region.objects.values_list('name', flat=True).get(pk=region.pk), 'Original') def test_partial_save_then_full_save_refreshes_path(self): + """ + Excluding parent must not advance _loaded_parent_id, otherwise a later + full save can leave the in-memory path stale. + """ r1 = Region.objects.create(name='R1', slug='r1-uf') r2 = Region.objects.create(name='R2', slug='r2-uf') obj = Region.objects.create(name='Obj', slug='obj-uf')