Jelajahi Sumber

Fixes #23321: Fix TypeError when creating a DeviceType with images via the REST API (#23330)

Avoid treating new uploads as stored image paths during device type
creation via the REST API. Read existing image names before saving and
delete replaced files after commit to preserve the originals on rollback.

Add regression coverage for image uploads, replacement, and cleanup.

Co-authored-by: Martin Hauser <mhauser@netboxlabs.com>
Jeremy Stretch 21 jam lalu
induk
melakukan
61fc5217e4

+ 23 - 11
netbox/dcim/models/devices.py

@@ -1,5 +1,5 @@
 import decimal
-from functools import cached_property
+from functools import cached_property, partial
 
 import yaml
 from django.contrib.contenttypes.fields import GenericForeignKey, GenericRelation
@@ -8,7 +8,7 @@ from django.contrib.postgres.indexes import GistIndex
 from django.core.exceptions import ValidationError
 from django.core.files.storage import default_storage
 from django.core.validators import MaxValueValidator, MinValueValidator
-from django.db import models
+from django.db import models, router, transaction
 from django.db.models import F, ProtectedError, prefetch_related_objects
 from django.db.models.functions import Lower
 from django.db.models.signals import post_save
@@ -27,6 +27,7 @@ from netbox.config import ConfigItem
 from netbox.models import NestedLtreeGroupModel, OrganizationalModel, PrimaryModel
 from netbox.models.features import ContactsMixin, 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.prefetch import get_prefetchable_fields
@@ -244,10 +245,6 @@ class DeviceType(ImageAttachmentsMixin, PrimaryModel, WeightMixin):
         # Save a copy of u_height for validation in clean()
         self._original_u_height = self.__dict__.get('u_height')
 
-        # Save references to the original front/rear images
-        self._original_front_image = self.__dict__.get('front_image')
-        self._original_rear_image = self.__dict__.get('rear_image')
-
     @property
     def full_name(self):
         return f"{self.manufacturer} {self.model}"
@@ -388,13 +385,28 @@ class DeviceType(ImageAttachmentsMixin, PrimaryModel, WeightMixin):
             })
 
     def save(self, *args, **kwargs):
+        update_fields = normalize_update_fields(kwargs)
+        # Use the write database for the image lookup and cleanup transaction.
+        using = kwargs.get('using') or router.db_for_write(self.__class__, instance=self)
+
+        # Retrieve stored image names only for fields being saved.
+        deferred_fields = self.get_deferred_fields()
+        image_fields = [
+            field_name for field_name in ('front_image', 'rear_image')
+            if field_name not in deferred_fields and (update_fields is None or field_name in update_fields)
+        ]
+        original_images = {}
+        if image_fields and self.pk is not None:
+            original_images = DeviceType.objects.using(using).filter(pk=self.pk).values(*image_fields).first() or {}
+
+        kwargs['using'] = using
         ret = super().save(*args, **kwargs)
 
-        # Delete any previously uploaded image files that are no longer in use
-        if self._original_front_image and self.front_image != self._original_front_image:
-            default_storage.delete(self._original_front_image)
-        if self._original_rear_image and self.rear_image != self._original_rear_image:
-            default_storage.delete(self._original_rear_image)
+        # Delete replaced files after commit to preserve them on rollback.
+        # Log cleanup failures without failing an already committed save.
+        for field_name, original_name in original_images.items():
+            if original_name and getattr(self, field_name).name != original_name:
+                transaction.on_commit(partial(default_storage.delete, original_name), using=using, robust=True)
 
         return ret
 

+ 52 - 0
netbox/dcim/tests/test_api.py

@@ -1,4 +1,6 @@
 import json
+import os
+import tempfile
 from unittest.mock import patch
 
 from django.conf import settings
@@ -31,6 +33,7 @@ from utilities.testing import (
     GraphQLFilterTest,
     GraphQLQueryTest,
     create_test_device,
+    create_test_image,
     create_test_nat_ip_pair,
     disable_logging,
     disable_warnings,
@@ -1577,6 +1580,55 @@ class DeviceTypeTestCase(APIViewTestCases.APIViewTestCase):
             },
         ]
 
+    def test_create_object_with_images(self):
+        """
+        Create a DeviceType with front/rear images via a multipart request.
+        """
+        self.add_permissions('dcim.add_devicetype')
+        # A nested field is read from "<field>.id" in a multipart request
+        data = {
+            'manufacturer.id': Manufacturer.objects.first().pk,
+            'model': 'Device Type 7',
+            'slug': 'device-type-7',
+            'front_image': create_test_image('front.png'),
+            'rear_image': create_test_image('rear.png'),
+        }
+
+        with tempfile.TemporaryDirectory() as media_root, override_settings(MEDIA_ROOT=media_root):
+            response = self.client.post(self._get_list_url(), data, format='multipart', **self.header)
+            self.assertHttpStatus(response, status.HTTP_201_CREATED)
+            device_type = DeviceType.objects.get(pk=response.data['id'])
+            self.assertTrue(os.path.exists(device_type.front_image.path))
+            self.assertTrue(os.path.exists(device_type.rear_image.path))
+
+    def test_update_object_image(self):
+        """
+        Replace a DeviceType's front image via a multipart request, which should delete the original file.
+        """
+        self.add_permissions('dcim.change_devicetype')
+
+        with tempfile.TemporaryDirectory() as media_root, override_settings(MEDIA_ROOT=media_root):
+            device_type = DeviceType.objects.first()
+            device_type.front_image = create_test_image('front.png')
+            device_type.rear_image = create_test_image('rear.png')
+            device_type.save()
+            front_image_path = device_type.front_image.path
+            rear_image_path = device_type.rear_image.path
+
+            data = {
+                'front_image': create_test_image('front2.png'),
+            }
+            with self.captureOnCommitCallbacks(execute=True):
+                response = self.client.patch(
+                    self._get_detail_url(device_type), data, format='multipart', **self.header
+                )
+            self.assertHttpStatus(response, status.HTTP_200_OK)
+            device_type.refresh_from_db()
+            self.assertFalse(os.path.exists(front_image_path))
+            self.assertTrue(os.path.exists(device_type.front_image.path))
+            self.assertNotEqual(device_type.front_image.path, front_image_path)
+            self.assertTrue(os.path.exists(rear_image_path))
+
 
 class ModuleTypeTestCase(APIViewTestCases.APIViewTestCase):
     model = ModuleType

+ 206 - 3
netbox/dcim/tests/test_models.py

@@ -1,12 +1,16 @@
+import os
+import tempfile
 import uuid
 from decimal import Decimal
+from unittest.mock import patch
 
 from django.core.exceptions import ValidationError
-from django.db import connection
+from django.db import connection, transaction
 from django.db.models import ProtectedError
 from django.db.models.signals import post_save
-from django.test import RequestFactory, TestCase, tag
+from django.test import RequestFactory, TestCase, override_settings, tag
 from django.test.utils import CaptureQueriesContext
+from django.utils.text import slugify
 
 from circuits.models import *
 from core.choices import ObjectChangeActionChoices
@@ -21,7 +25,7 @@ from netbox.context_managers import event_tracking
 from tenancy.models import Tenant
 from users.models import User
 from utilities.data import drange
-from utilities.testing import create_test_device
+from utilities.testing import create_test_device, create_test_image
 from virtualization.models import Cluster, ClusterType
 
 
@@ -221,6 +225,205 @@ class DeviceTypeTestCase(TestCase):
         self.assertEqual(device_type.interface_template_count, 1)
 
 
+class DeviceTypeImageTestCase(TestCase):
+    """
+    Test the handling of DeviceType front/rear image files.
+    """
+    def setUp(self):
+        media_root = self.enterContext(tempfile.TemporaryDirectory())
+        self.enterContext(override_settings(MEDIA_ROOT=media_root))
+        self.manufacturer = Manufacturer.objects.create(name='Manufacturer 1', slug='manufacturer-1')
+
+    def _create_device_type(self, model='Device Type 1', **images):
+        """
+        Create a DeviceType with an uploaded image for each image field mapped to a filename.
+        """
+        return DeviceType.objects.create(
+            manufacturer=self.manufacturer,
+            model=model,
+            slug=slugify(model),
+            **{field_name: create_test_image(filename) for field_name, filename in images.items()},
+        )
+
+    def _save(self, device_type, **kwargs):
+        # Execute the on-commit deletion of replaced image files, which TestCase otherwise discards
+        with self.captureOnCommitCallbacks(execute=True):
+            device_type.save(**kwargs)
+
+    def test_create_with_images(self):
+        device_type = self._create_device_type(front_image='front.png', rear_image='rear.png')
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+        self.assertTrue(os.path.exists(device_type.rear_image.path))
+
+    def test_replace_and_clear_images(self):
+        device_type = self._create_device_type(front_image='front1.png', rear_image='rear.png')
+        front_image_path = device_type.front_image.path
+        rear_image_path = device_type.rear_image.path
+
+        # Replacing an image should delete the original file
+        device_type = DeviceType.objects.get(pk=device_type.pk)
+        device_type.front_image = create_test_image('front2.png')
+        self._save(device_type)
+        self.assertFalse(os.path.exists(front_image_path))
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+
+        # Clearing an image should delete its file
+        device_type.rear_image = None
+        self._save(device_type)
+        self.assertFalse(os.path.exists(rear_image_path))
+
+    def test_replace_image_on_same_instance(self):
+        """
+        Replacing an image repeatedly on the instance it was created with should delete each prior file.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        first_image_path = device_type.front_image.path
+
+        device_type.front_image = create_test_image('front2.png')
+        self._save(device_type)
+        second_image_path = device_type.front_image.path
+        self.assertFalse(os.path.exists(first_image_path))
+        self.assertTrue(os.path.exists(second_image_path))
+
+        device_type.front_image = create_test_image('front3.png')
+        self._save(device_type)
+        self.assertFalse(os.path.exists(second_image_path))
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+
+    def test_replace_initial_image_path_before_first_save(self):
+        """
+        Replacing an image path assigned to a new DeviceType before its first save should not delete that file, which
+        here belongs to another DeviceType.
+        """
+        device_type1 = self._create_device_type(front_image='front1.png')
+        shared_image_path = device_type1.front_image.path
+
+        device_type2 = DeviceType(
+            manufacturer=self.manufacturer,
+            model='Device Type 2',
+            slug='device-type-2',
+            front_image=device_type1.front_image.name,
+        )
+        device_type2.front_image = create_test_image('front2.png')
+        self._save(device_type2)
+        self.assertTrue(os.path.exists(shared_image_path))
+        self.assertTrue(os.path.exists(device_type2.front_image.path))
+
+    def test_update_fields_omitting_image(self):
+        """
+        Saving with update_fields which omits an image field should not delete its stored file, which a later full
+        save then replaces.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        original_image_path = device_type.front_image.path
+
+        device_type = DeviceType.objects.get(pk=device_type.pk)
+        device_type.front_image = create_test_image('front2.png')
+        device_type.description = 'New description'
+        self._save(device_type, update_fields=['description'])
+        self.assertTrue(os.path.exists(original_image_path))
+        device_type.refresh_from_db(fields=['front_image'])
+        self.assertEqual(device_type.front_image.path, original_image_path)
+
+        device_type.front_image = create_test_image('front2.png')
+        self._save(device_type)
+        self.assertFalse(os.path.exists(original_image_path))
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+
+    def test_update_fields_iterator(self):
+        """
+        Replacing an image via a save with iterator-valued update_fields should delete the original file.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        original_image_path = device_type.front_image.path
+
+        device_type = DeviceType.objects.get(pk=device_type.pk)
+        device_type.front_image = create_test_image('front2.png')
+        self._save(device_type, update_fields=iter(['front_image']))
+        self.assertFalse(os.path.exists(original_image_path))
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+
+    def test_deferred_images(self):
+        """
+        Saving a DeviceType with deferred image fields should neither load them nor delete their stored files.
+        """
+        device_type = self._create_device_type(front_image='front.png', rear_image='rear.png')
+        front_image_path = device_type.front_image.path
+        rear_image_path = device_type.rear_image.path
+
+        device_type = DeviceType.objects.defer('front_image', 'rear_image').get(pk=device_type.pk)
+        device_type.description = 'New description'
+        self._save(device_type)
+        self.assertTrue({'front_image', 'rear_image'}.issubset(device_type.get_deferred_fields()))
+        self.assertTrue(os.path.exists(front_image_path))
+        self.assertTrue(os.path.exists(rear_image_path))
+
+    def test_replace_deferred_image(self):
+        """
+        Replacing an image which was deferred when the DeviceType was loaded should delete the original file.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        original_image_path = device_type.front_image.path
+
+        device_type = DeviceType.objects.defer('front_image').get(pk=device_type.pk)
+        device_type.front_image = create_test_image('front2.png')
+        self._save(device_type)
+        self.assertFalse(os.path.exists(original_image_path))
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+
+    def test_replace_image_on_stale_instance(self):
+        """
+        Replacing an image on an instance loaded before another change to that image should delete the file currently
+        stored, not the one originally loaded.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        stale_instance = DeviceType.objects.get(pk=device_type.pk)
+
+        device_type.front_image = create_test_image('front2.png')
+        self._save(device_type)
+        second_image_path = device_type.front_image.path
+
+        stale_instance.front_image = create_test_image('front3.png')
+        self._save(stale_instance)
+        self.assertFalse(os.path.exists(second_image_path))
+        self.assertTrue(os.path.exists(stale_instance.front_image.path))
+
+    def test_replace_image_rolled_back(self):
+        """
+        Replacing an image within a transaction which is rolled back should not delete the original file.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        original_image_path = device_type.front_image.path
+
+        device_type.front_image = create_test_image('front2.png')
+        with self.captureOnCommitCallbacks(execute=True), self.assertRaises(RuntimeError):
+            with transaction.atomic():
+                device_type.save()
+                raise RuntimeError('Roll back the transaction')
+        self.assertTrue(os.path.exists(original_image_path))
+        device_type.refresh_from_db()
+        self.assertEqual(device_type.front_image.path, original_image_path)
+
+    def test_replace_image_storage_error(self):
+        """
+        A failure to delete a replaced image file should be logged rather than raised, as the change itself has
+        already been committed.
+        """
+        device_type = self._create_device_type(front_image='front1.png')
+        original_image_name = device_type.front_image.name
+
+        # A non-robust on-commit callback would re-raise the OSError here
+        device_type.front_image = create_test_image('front2.png')
+        with patch(
+            'dcim.models.devices.default_storage.delete', side_effect=OSError('Storage unavailable')
+        ) as mock_delete:
+            self._save(device_type)
+        mock_delete.assert_called_once_with(original_image_name)
+        device_type.refresh_from_db()
+        self.assertTrue(os.path.exists(device_type.front_image.path))
+        self.assertTrue(device_type.front_image.name.endswith('front2.png'))
+
+
 class ModuleTypeTestCase(TestCase):
 
     def test_component_template_counts(self):

+ 12 - 0
netbox/utilities/testing/utils.py

@@ -1,3 +1,4 @@
+import io
 import json
 import logging
 import random
@@ -6,7 +7,9 @@ import string
 from contextlib import contextmanager
 
 from django.contrib.auth.models import Permission
+from django.core.files.uploadedfile import SimpleUploadedFile
 from django.utils.text import slugify
+from PIL import Image
 
 from core.models import ObjectType
 from dcim.models import Device, DeviceRole, DeviceType, Manufacturer, Site
@@ -55,6 +58,15 @@ def create_test_device(name, site=None, **attrs):
     return device
 
 
+def create_test_image(filename):
+    """
+    Convenience method for creating an uploaded PNG image file.
+    """
+    image = io.BytesIO()
+    Image.new('RGB', (1, 1)).save(image, format='PNG')
+    return SimpleUploadedFile(name=filename, content=image.getvalue(), content_type='image/png')
+
+
 def create_test_virtualmachine(name):
     """
     Convenience method for creating a VirtualMachine.